diff --git a/mjx/mujoco/mjx/_src/io.py b/mjx/mujoco/mjx/_src/io.py index 8feef14d..4199936a 100644 --- a/mjx/mujoco/mjx/_src/io.py +++ b/mjx/mujoco/mjx/_src/io.py @@ -1943,6 +1943,7 @@ def set_state( def create_render_context( mjm: mujoco.MjModel, nworld: int, + devices: Optional[Sequence[str]] = None, **kwargs, ): """Creates a render context. @@ -1953,6 +1954,9 @@ def create_render_context( because Warp creates arrays of size nworld that are not exposed to JAX. Thus we cannot use JAX transforms like vmap with the render context. + devices: optional list of device names (e.g. ['cuda:0', 'cuda:1']). + If provided, rendering workloads are sharded across these devices. + By default, devices is None and the default device from wp.get_device(None) is used. **kwargs: forwarded to the render context constructor. Returns: @@ -1960,4 +1964,6 @@ def create_render_context( """ _check_warp_installed() from mujoco.mjx.warp import io as mjxw_io # pylint: disable=g-import-not-at-top # pytype: disable=import-error - return mjxw_io.create_render_context(mjm, nworld=nworld, **kwargs) + return mjxw_io.create_render_context( + mjm, nworld=nworld, devices=devices, **kwargs + ) diff --git a/mjx/mujoco/mjx/_src/render_util.py b/mjx/mujoco/mjx/_src/render_util.py index 58e98ef4..cfe6bf0b 100644 --- a/mjx/mujoco/mjx/_src/render_util.py +++ b/mjx/mujoco/mjx/_src/render_util.py @@ -44,7 +44,7 @@ def get_rgb( else: raise RuntimeError('Warp not installed.') - warp_rc = mjxw_render._MJX_RENDER_CONTEXT_BUFFERS[rc.key] + warp_rc = mjxw_render._MJX_RENDER_CONTEXT_BUFFERS[(rc.key, None)] 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]) @@ -84,7 +84,7 @@ def get_depth( import mujoco.mjx.warp.render as mjxw_render # pylint: disable=g-import-not-at-top # pytype: disable=import-error else: raise RuntimeError('Warp not installed.') - warp_rc = mjxw_render._MJX_RENDER_CONTEXT_BUFFERS[rc.key] + warp_rc = mjxw_render._MJX_RENDER_CONTEXT_BUFFERS[(rc.key, None)] 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]) diff --git a/mjx/mujoco/mjx/_src/render_util_test.py b/mjx/mujoco/mjx/_src/render_util_test.py index e62658ac..ed82f5c0 100644 --- a/mjx/mujoco/mjx/_src/render_util_test.py +++ b/mjx/mujoco/mjx/_src/render_util_test.py @@ -56,7 +56,7 @@ class RenderUtilTest(absltest.TestCase): with mock.patch.dict( 'mujoco.mjx.warp.render._MJX_RENDER_CONTEXT_BUFFERS', - {0: warp_rc}, + {(0, None): warp_rc}, ): rgb = jax.jit(render_util.get_rgb, static_argnums=(0, 1))(rc, 0, rgb_data) @@ -70,7 +70,7 @@ class RenderUtilTest(absltest.TestCase): with mock.patch.dict( 'mujoco.mjx.warp.render._MJX_RENDER_CONTEXT_BUFFERS', - {0: warp_rc}, + {(0, None): warp_rc}, ): rgb = jax.jit( jax.vmap(render_util.get_rgb, in_axes=(None, None, 0)), @@ -87,7 +87,7 @@ class RenderUtilTest(absltest.TestCase): with mock.patch.dict( 'mujoco.mjx.warp.render._MJX_RENDER_CONTEXT_BUFFERS', - {0: warp_rc}, + {(0, None): warp_rc}, ): depth = jax.jit(render_util.get_depth, static_argnums=(0, 1, 3))( rc, 0, depth_data, 5.0 @@ -103,7 +103,7 @@ class RenderUtilTest(absltest.TestCase): with mock.patch.dict( 'mujoco.mjx.warp.render._MJX_RENDER_CONTEXT_BUFFERS', - {0: warp_rc}, + {(0, None): warp_rc}, ): depth = jax.jit( jax.vmap(render_util.get_depth, in_axes=(None, None, 0, None)), diff --git a/mjx/mujoco/mjx/warp/bvh.py b/mjx/mujoco/mjx/warp/bvh.py index 573c24e4..a86bad63 100644 --- a/mjx/mujoco/mjx/warp/bvh.py +++ b/mjx/mujoco/mjx/warp/bvh.py @@ -88,7 +88,7 @@ def _refit_bvh_shim( _d.geom_xmat = geom_xmat _d.geom_xpos = geom_xpos _d.nworld = nworld - render_context = _MJX_RENDER_CONTEXT_BUFFERS[rc_id] + render_context = _MJX_RENDER_CONTEXT_BUFFERS[(rc_id, wp.get_device().ordinal)] dummy.zero_() mjwarp.refit_bvh(_m, _d, render_context) diff --git a/mjx/mujoco/mjx/warp/collision_driver.py b/mjx/mujoco/mjx/warp/collision_driver.py index 35a85720..aa88f118 100644 --- a/mjx/mujoco/mjx/warp/collision_driver.py +++ b/mjx/mujoco/mjx/warp/collision_driver.py @@ -44,6 +44,7 @@ _e = mjwarp.Constraint( **{f.name: None for f in dataclasses.fields(mjwarp.Constraint) if f.init} ) + @ffi.format_args_for_warp def _collision_shim( # Model diff --git a/mjx/mujoco/mjx/warp/forward.py b/mjx/mujoco/mjx/warp/forward.py index 349e5ed2..db52a8cd 100644 --- a/mjx/mujoco/mjx/warp/forward.py +++ b/mjx/mujoco/mjx/warp/forward.py @@ -44,6 +44,7 @@ _e = mjwarp.Constraint( **{f.name: None for f in dataclasses.fields(mjwarp.Constraint) if f.init} ) + @ffi.format_args_for_warp def _forward_shim( # Model diff --git a/mjx/mujoco/mjx/warp/io.py b/mjx/mujoco/mjx/warp/io.py index eb9af42c..db0f2df9 100644 --- a/mjx/mujoco/mjx/warp/io.py +++ b/mjx/mujoco/mjx/warp/io.py @@ -19,26 +19,43 @@ import threading import mujoco from mujoco.mjx.warp.types import RenderContext import mujoco.mjx.third_party.mujoco_warp as mjw +import warp as wp _MJX_RENDER_CONTEXT_COUNTER = 0 _MJX_RENDER_CONTEXT_LOCK = threading.Lock() _MJX_RENDER_CONTEXT_BUFFERS = {} +def _create_context(mjm, nworld, device, **kwargs): + with wp.ScopedDevice(device): + 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.rgb_data = None + ctx.depth_data = None + return ctx + + def create_render_context( mjm: mujoco.MjModel, nworld: int, + devices: list[str | None] | None = None, **kwargs, ): - rc = mjw.create_render_context(mjm=mjm, nworld=nworld, **kwargs) - rc.rgb_data_shape = rc.rgb_data.shape - rc.depth_data_shape = rc.depth_data.shape - rc.rgb_data = None - rc.depth_data = None - global _MJX_RENDER_CONTEXT_COUNTER + + if not devices: + devices = [None] + + contexts = [_create_context(mjm, nworld, d, **kwargs) for d in devices] + with _MJX_RENDER_CONTEXT_LOCK: _MJX_RENDER_CONTEXT_COUNTER += 1 key = _MJX_RENDER_CONTEXT_COUNTER - _MJX_RENDER_CONTEXT_BUFFERS[key] = rc + for d, ctx in zip(devices, contexts): + ordinal = wp.get_device(d).ordinal + _MJX_RENDER_CONTEXT_BUFFERS[(key, ordinal)] = ctx + if (key, None) not in _MJX_RENDER_CONTEXT_BUFFERS: + # save the first context as the default context + _MJX_RENDER_CONTEXT_BUFFERS[(key, None)] = contexts[0] return RenderContext(key, _owner=True) diff --git a/mjx/mujoco/mjx/warp/render.py b/mjx/mujoco/mjx/warp/render.py index 091f72d0..73ced206 100644 --- a/mjx/mujoco/mjx/warp/render.py +++ b/mjx/mujoco/mjx/warp/render.py @@ -110,14 +110,14 @@ def _render_shim( _d.light_xdir = light_xdir _d.light_xpos = light_xpos _d.nworld = nworld - render_context = _MJX_RENDER_CONTEXT_BUFFERS[rc_id] + render_context = _MJX_RENDER_CONTEXT_BUFFERS[(rc_id, wp.get_device().ordinal)] render_context.rgb_data = rgb render_context.depth_data = depth mjwarp.render(_m, _d, render_context) def _render_jax_impl(m: types.Model, d: types.Data, ctx: RenderContext): - render_ctx = _MJX_RENDER_CONTEXT_BUFFERS[ctx.key] + render_ctx = _MJX_RENDER_CONTEXT_BUFFERS[(ctx.key, None)] output_dims = { 'rgb': render_ctx.rgb_data_shape, 'depth': render_ctx.depth_data_shape, diff --git a/mjx/mujoco/mjx/warp/smooth.py b/mjx/mujoco/mjx/warp/smooth.py index 217cf873..6e1c0ea8 100644 --- a/mjx/mujoco/mjx/warp/smooth.py +++ b/mjx/mujoco/mjx/warp/smooth.py @@ -44,6 +44,7 @@ _e = mjwarp.Constraint( **{f.name: None for f in dataclasses.fields(mjwarp.Constraint) if f.init} ) + @ffi.format_args_for_warp def _kinematics_shim( # Model diff --git a/mjx/mujoco/mjx/warp/types.py b/mjx/mujoco/mjx/warp/types.py index d39d4861..03d403a4 100644 --- a/mjx/mujoco/mjx/warp/types.py +++ b/mjx/mujoco/mjx/warp/types.py @@ -126,7 +126,11 @@ class RenderContext: if lock is None or buffers is None: return with lock: - buffers.pop(self.key, None) + keys_to_remove = [ + k for k in buffers.keys() if isinstance(k, tuple) and k[0] == self.key + ] + for k in keys_to_remove: + buffers.pop(k, None) class StatisticWarp(PyTreeNode): diff --git a/mjx/mujoco/mjx/warp/visualize_render.py b/mjx/mujoco/mjx/warp/visualize_render.py index 42ed02f1..6a3b5afa 100644 --- a/mjx/mujoco/mjx/warp/visualize_render.py +++ b/mjx/mujoco/mjx/warp/visualize_render.py @@ -57,6 +57,9 @@ _WP_KERNEL_CACHE_DIR = flags.DEFINE_string( '/tmp/wp_kernel_cache_dir_visualize_render', 'warp kernel cache directory', ) +_PMAP = flags.DEFINE_boolean( + 'pmap', False, 'also render with pmap across GPUs and compare' +) _COMPILER_OPTIONS = {'xla_gpu_graph_min_graph_size': 1} jax_jit = functools.partial(jax.jit, compiler_options=_COMPILER_OPTIONS) @@ -107,6 +110,7 @@ def _main(_: Sequence[str]): print(f' camera_id : {_CAMERA_ID.value}') print(f' use_textures: {_USE_TEXTURES.value}') print(f' use_shadows : {_USE_SHADOWS.value}') + print(f' pmap : {_PMAP.value}') print(f' output_dir : {_OUTPUT_DIR.value}\n') mx = mjx.put_model(m, impl='warp') @@ -143,7 +147,6 @@ def _main(_: Sequence[str]): enabled_geom_groups=[0, 1, 2], ) - print('rendering...') dx_batch = jax_jit(jax.vmap(bvh.refit_bvh, in_axes=(None, 0, None)))( mx, dx_batch, rc ) @@ -187,6 +190,62 @@ def _main(_: Sequence[str]): ) _save_tiled(depth_rgb, depth_tiled_path) + if _PMAP.value: + ndevices = jax.local_device_count() + nworld = _NWORLD.value + nworld_per_device = nworld // ndevices + assert nworld >= ndevices and nworld % ndevices == 0, ( + f'--pmap requires nworld ({nworld}) divisible by device count' + f' ({ndevices})' + ) + print(f'\nrendering (pmap across {ndevices} devices)...') + + device_strs = [f'cuda:{i}' for i in range(ndevices)] + + pmap_rc = io.create_render_context( + mjm=m, + nworld=nworld_per_device, + devices=device_strs, + cam_res=(_WIDTH.value, _HEIGHT.value), + use_textures=_USE_TEXTURES.value, + use_shadows=_USE_SHADOWS.value, + render_rgb=True, + render_depth=True, + enabled_geom_groups=[0, 1, 2], + ) + + devices = jax.local_devices()[:ndevices] + mesh = jax.sharding.Mesh(np.array(devices), axis_names=('i',)) + P = jax.sharding.PartitionSpec + sharded = jax.sharding.NamedSharding(mesh, P('i')) + + def safe_shard(x, sharding): + # Go through CPU to avoid P2P DMA issues on certain machines. + x_cpu = jax.device_put(x, jax.devices('cpu')[0]) + if x_cpu.ndim > 0 and x_cpu.shape[0] == nworld: + reshaped = x_cpu.reshape(ndevices, nworld_per_device, *x_cpu.shape[1:]) + else: + reshaped = jp.stack([x_cpu] * ndevices) + return jax.device_put(reshaped, sharding) + + dx_pmap = jax.tree.map(lambda x: safe_shard(x, sharded), dx_batch) + mx_pmap = jax.tree.map(lambda x: safe_shard(x, sharded), mx) + + def inner(mx, dx): + dx = bvh.refit_bvh(mx, dx, pmap_rc) + out = render.render(mx, dx, pmap_rc) + return render_util.get_rgb(pmap_rc, _CAMERA_ID.value, out[0]) + + inner = jax.vmap(inner, in_axes=(None, 0)) + out = jax.pmap(inner)(mx_pmap, dx_pmap) + + pmap_rgb = jax.device_put(out, jax.devices('cpu')[0]).reshape(-1, *out.shape[2:]) + + pmap_tiled_path = os.path.join( + _OUTPUT_DIR.value, f'pmap_tiled_{_CAMERA_ID.value}.png' + ) + _save_tiled(pmap_rgb, pmap_tiled_path) + print('\ndone.')