Merge pull request #3381 from tkelestemur:tarik/mjx-warp-put-model-nworld
PiperOrigin-RevId: 945578693 Change-Id: I873e4f47551d02f506d3af37040e1014fa0edb61
This commit is contained in:
@@ -442,11 +442,13 @@ def _put_model_warp(
|
||||
m: mujoco.MjModel,
|
||||
graph_mode: mjxw.types.GraphMode,
|
||||
device: Optional[jax.Device] = None,
|
||||
batch_sizes: Optional[Dict[str, int]] = None,
|
||||
) -> types.Model:
|
||||
"""Puts mujoco.MjModel onto a device, resulting in mjx.Model."""
|
||||
with wp.ScopedDevice('cpu'): # pylint: disable=undefined-variable
|
||||
mw = mjwp.put_model(m) # pylint: disable=undefined-variable
|
||||
mw = mjwp.put_model(m, batch_sizes=batch_sizes) # pylint: disable=undefined-variable
|
||||
|
||||
batch_sizes = batch_sizes or {}
|
||||
fields = {f.name for f in types.Model.fields() if f.name != '_impl'}
|
||||
fields = {f: getattr(m, f) for f in fields}
|
||||
# Grab MJW private Option fields, and assume that public MjOption fields are
|
||||
@@ -467,7 +469,9 @@ def _put_model_warp(
|
||||
if not hasattr(mw, k) or k in ('stat', 'opt'):
|
||||
continue
|
||||
field = _wp_to_np_type(getattr(mw, k), k)
|
||||
if mjxw.types._BATCH_DIM['Model'].get(k, False): # pylint: disable=protected-access
|
||||
if ( # pylint: disable=protected-access
|
||||
k not in batch_sizes and mjxw.types._BATCH_DIM['Model'].get(k, False)
|
||||
):
|
||||
field = field.reshape(field.shape[1:])
|
||||
if k == 'geom_dataid' and field.ndim > 1:
|
||||
# Batched geom_dataid is not supported in MJX.
|
||||
@@ -477,7 +481,9 @@ def _put_model_warp(
|
||||
impl_fields = {}
|
||||
for k in mjxw.types.ModelWarp.__annotations__.keys():
|
||||
field = _wp_to_np_type(getattr(mw, k), k)
|
||||
if mjxw.types._BATCH_DIM['Model'].get(k, False): # pylint: disable=protected-access
|
||||
if ( # pylint: disable=protected-access
|
||||
k not in batch_sizes and mjxw.types._BATCH_DIM['Model'].get(k, False)
|
||||
):
|
||||
field = field.reshape(field.shape[1:])
|
||||
impl_fields[k] = field
|
||||
|
||||
@@ -534,6 +540,7 @@ def put_model(
|
||||
impl: Optional[Union[str, types.Impl]] = None,
|
||||
graph_mode: Optional[mjxw.types.GraphMode] = None,
|
||||
keepalive_refs: Optional[Dict[int, Any]] = None,
|
||||
batch_sizes: Optional[Dict[str, int]] = None,
|
||||
) -> types.Model:
|
||||
"""Puts mujoco.MjModel onto a device, resulting in mjx.Model.
|
||||
|
||||
@@ -546,6 +553,7 @@ def put_model(
|
||||
keepalive_refs: optional dict to store references to underlying MuJoCo
|
||||
objects, preventing them from being garbage collected. Required for CPP
|
||||
impl to keep the model alive.
|
||||
batch_sizes: optional per-field leading batch sizes for Warp model fields.
|
||||
|
||||
Returns:
|
||||
an mjx.Model placed on device
|
||||
@@ -561,7 +569,7 @@ def put_model(
|
||||
elif impl == types.Impl.WARP:
|
||||
_check_warp_installed()
|
||||
graph_mode = graph_mode or getattr(mjxw.types.GraphMode, 'WARP')
|
||||
return _put_model_warp(m, graph_mode, device)
|
||||
return _put_model_warp(m, graph_mode, device, batch_sizes=batch_sizes)
|
||||
elif impl == types.Impl.CPP:
|
||||
return _put_model_cpp(m, device, keepalive_refs=keepalive_refs)
|
||||
else:
|
||||
|
||||
@@ -352,6 +352,43 @@ class ModelIOTest(parameterized.TestCase):
|
||||
|
||||
_ = jax.tree.map_with_path(check_ndim, mx)
|
||||
|
||||
def test_put_model_warp_batch_sizes(self):
|
||||
"""Tests put_model can add an nworld axis to selected Warp model fields."""
|
||||
if not mjxw.WARP_INSTALLED:
|
||||
self.skipTest('Warp not installed.')
|
||||
if not hasattr(mjxw_types.GraphMode, 'WARP'):
|
||||
self.skipTest('Warp JAX FFI graph modes unavailable.')
|
||||
|
||||
m = mujoco.MjModel.from_xml_string("""
|
||||
<mujoco>
|
||||
<asset>
|
||||
<texture name="red" type="2d" builtin="flat" width="4" height="4"
|
||||
rgb1="1 0 0" rgb2="1 0 0"/>
|
||||
<material name="mat" texture="red" rgba="0.5 0.6 0.7 1"/>
|
||||
</asset>
|
||||
<worldbody>
|
||||
<geom type="sphere" size="0.1" material="mat"/>
|
||||
</worldbody>
|
||||
</mujoco>
|
||||
""")
|
||||
|
||||
nworld = 3
|
||||
mx = mjx.put_model(m, impl='warp', batch_sizes={'mat_texid': nworld})
|
||||
|
||||
self.assertEqual(mx.mat_texid.shape, (nworld,) + m.mat_texid.shape)
|
||||
self.assertEqual(mx.geom_pos.shape, (m.ngeom, 3))
|
||||
self.assertEqual(mx.geom_size.shape, (m.ngeom, 3))
|
||||
self.assertEqual(mx.mat_rgba.shape, (m.nmat, 4))
|
||||
self.assertEqual(mx.geom_type.shape, (m.ngeom,))
|
||||
self.assertEqual(mx.geom_dataid.shape, (m.ngeom,))
|
||||
|
||||
np.testing.assert_array_equal(
|
||||
np.asarray(mx.mat_texid), np.repeat(m.mat_texid[None], nworld, axis=0)
|
||||
)
|
||||
np.testing.assert_allclose(np.asarray(mx.geom_pos), m.geom_pos)
|
||||
np.testing.assert_allclose(np.asarray(mx.geom_size), m.geom_size)
|
||||
np.testing.assert_allclose(np.asarray(mx.mat_rgba), m.mat_rgba)
|
||||
|
||||
@parameterized.parameters('JAX', 'WARP', None)
|
||||
def test_put_model_warp_graph_mode(self, mode: str | None):
|
||||
"""Tests that put_model accepts graph_mode parameter."""
|
||||
|
||||
Reference in New Issue
Block a user