Add pmap support for mjx-warp render.

PiperOrigin-RevId: 872533248
Change-Id: I828303866c31d33d2367dd9fba4eb42028c67540
This commit is contained in:
Baruch Tabanpour
2026-02-19 13:13:41 -08:00
committed by Copybara-Service
parent 84431fbf7e
commit 62a32386d6
11 changed files with 108 additions and 19 deletions
+7 -1
View File
@@ -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
)
+2 -2
View File
@@ -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])
+4 -4
View File
@@ -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)),
+1 -1
View File
@@ -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)
+1
View File
@@ -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
+1
View File
@@ -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
+24 -7
View File
@@ -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)
+2 -2
View File
@@ -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,
+1
View File
@@ -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
+5 -1
View File
@@ -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):
+60 -1
View File
@@ -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.')