From 0be7e6a1f6c32c5707d9d78ee2831bed6e9a78f0 Mon Sep 17 00:00:00 2001 From: Baruch Tabanpour Date: Mon, 6 Jan 2025 10:27:16 -0800 Subject: [PATCH] Fix #2306. PiperOrigin-RevId: 712575320 Change-Id: Ia74ef0b31b4e7098647340760572a1f6368eab64 --- mjx/mujoco/mjx/_src/io.py | 3 ++- mjx/mujoco/mjx/_src/io_test.py | 9 +++++++++ 2 files changed, 11 insertions(+), 1 deletion(-) 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)