From b4135b571f8e8ec857e124092e5ce1d6b6be399c Mon Sep 17 00:00:00 2001 From: Mustafa Hasekioglu Date: Sat, 6 Jun 2026 16:15:23 -0400 Subject: [PATCH] fix island disable fields dimension mismatch --- mjx/mujoco/mjx/_src/io.py | 4 ++++ mjx/mujoco/mjx/_src/io_test.py | 20 +++++++++++++++++++- 2 files changed, 23 insertions(+), 1 deletion(-) diff --git a/mjx/mujoco/mjx/_src/io.py b/mjx/mujoco/mjx/_src/io.py index 70d7f973..03c7af65 100644 --- a/mjx/mujoco/mjx/_src/io.py +++ b/mjx/mujoco/mjx/_src/io.py @@ -1307,6 +1307,10 @@ def _get_data_into_warp( if isinstance(value, np.ndarray) and value.shape: result_field = getattr(result_i, field.name) if result_field.shape != value.shape: + # When ENABLE_ISLANDS is False, mujoco_warp allocates island fields with width 0, + # while the host MjData sizes to nv/ntree. Skip to prevent mismatch. + if value.size == 0 and not mjxw.mjwp_io.ENABLE_ISLANDS: + continue raise ValueError( f'Input field {field.name} has shape {value.shape}, but output' f' has shape {result_field.shape}' diff --git a/mjx/mujoco/mjx/_src/io_test.py b/mjx/mujoco/mjx/_src/io_test.py index 953b1f72..4236adfc 100644 --- a/mjx/mujoco/mjx/_src/io_test.py +++ b/mjx/mujoco/mjx/_src/io_test.py @@ -797,12 +797,30 @@ class DataIOTest(parameterized.TestCase): if not mjx_io.has_cuda_gpu_device(): self.skipTest('No CUDA GPU device.') - m = mujoco.MjModel.from_xml_string('') + # Use a model with at least one kinematic tree so that island + # fields (e.g. dof_island, tree_island) are populated on the host + # MjData. + m = mujoco.MjModel.from_xml_string(""" + + + + + + + + + """) d = mujoco.MjData(m) mx = mjx.put_model(m, impl='warp') dx = mjx.make_data(m, impl='warp') mjx.get_data_into(d, mx, dx) + # Island data is not populated when ENABLE_ISLANDS is False, so the host + # MjData island fields should be left at their default (nv,)/(ntree,) + # shapes rather than triggering a shape mismatch. + self.assertEqual(d.dof_island.shape, (m.nv,)) + self.assertEqual(d.tree_island.shape, (m.ntree,)) + @parameterized.parameters(('jax',)) def test_get_data_into_wrong_shape(self, impl): """Tests that get_data_into throwsif input and output shapes don't match."""