diff --git a/mjx/mujoco/mjx/_src/io.py b/mjx/mujoco/mjx/_src/io.py index 2d90eb33..e5ae79c9 100644 --- a/mjx/mujoco/mjx/_src/io.py +++ b/mjx/mujoco/mjx/_src/io.py @@ -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( diff --git a/mjx/mujoco/mjx/_src/io_test.py b/mjx/mujoco/mjx/_src/io_test.py index 3fa5c7b6..a436d4ac 100644 --- a/mjx/mujoco/mjx/_src/io_test.py +++ b/mjx/mujoco/mjx/_src/io_test.py @@ -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)