Add pmap support for mjx-warp render.
PiperOrigin-RevId: 872533248 Change-Id: I828303866c31d33d2367dd9fba4eb42028c67540
This commit is contained in:
committed by
Copybara-Service
parent
84431fbf7e
commit
62a32386d6
@@ -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
|
||||
)
|
||||
|
||||
@@ -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])
|
||||
|
||||
@@ -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)),
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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.')
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user