Strip weak_types in mjx.put_model.

PiperOrigin-RevId: 677859138
Change-Id: I9ef010c0c37d04f9931adf884ab1dce0521b8567
This commit is contained in:
Google DeepMind
2024-09-23 10:53:43 -07:00
committed by Copybara-Service
parent 84e0e2839d
commit bbe8a0c516
2 changed files with 15 additions and 1 deletions
+10 -1
View File
@@ -29,6 +29,14 @@ import numpy as np
import scipy
def _strip_weak_type(tree):
def f(leaf):
if isinstance(leaf, jax.Array):
return leaf.astype(jax.dtypes.canonicalize_dtype(leaf.dtype))
return leaf
return jax.tree_util.tree_map(f, tree)
def _make_option(o: mujoco.MjOption) -> types.Option:
"""Returns mjx.Option given mujoco.MjOption."""
if o.integrator not in set(types.IntegratorType):
@@ -148,7 +156,8 @@ def put_model(
model = types.Model(**{k: copy.copy(v) for k, v in fields.items()})
return jax.device_put(model, device=device)
model = jax.device_put(model, device=device)
return _strip_weak_type(model)
def make_data(
+5
View File
@@ -88,6 +88,11 @@ class ModelIOTest(parameterized.TestCase):
def test_put_model(self):
m = mujoco.MjModel.from_xml_string(_MULTIPLE_CONVEX_OBJECTS)
mx = mjx.put_model(m)
def assert_not_weak_type(x):
if isinstance(x, jax.Array):
assert not x.weak_type
jax.tree_util.tree_map(assert_not_weak_type, mx)
self.assertEqual(mx.nq, m.nq)
self.assertEqual(mx.nv, m.nv)
self.assertEqual(mx.nu, m.nu)