diff --git a/mjx/mujoco/mjx/_src/io.py b/mjx/mujoco/mjx/_src/io.py index 55b6c08c..7a000a15 100644 --- a/mjx/mujoco/mjx/_src/io.py +++ b/mjx/mujoco/mjx/_src/io.py @@ -681,4 +681,5 @@ def put_data( # copy because device_put is async: data = types.Data(**{k: copy.copy(v) for k, v in fields.items()}) - return jax.device_put(data, device=device) + data = jax.device_put(data, device=device) + return _strip_weak_type(data) diff --git a/mjx/mujoco/mjx/_src/io_test.py b/mjx/mujoco/mjx/_src/io_test.py index 8950c677..b1730f33 100644 --- a/mjx/mujoco/mjx/_src/io_test.py +++ b/mjx/mujoco/mjx/_src/io_test.py @@ -310,6 +310,15 @@ class DataIOTest(parameterized.TestCase): np.testing.assert_allclose(dx.cvel, d.cvel) np.testing.assert_allclose(dx.cdof_dot, d.cdof_dot) + # check that there are no weak types + self.assertFalse( + any( + jax.tree_util.tree_flatten( + jax.tree_util.tree_map(lambda x: x.weak_type, dx) + )[0] + ) + ) + # check that qM is transformed properly qm = np.zeros((m.nv, m.nv), dtype=np.float64) mujoco.mj_fullM(m, qm, d.qM)