Merge pull request #3323 from mhasek:mhasek/fix_mjwarp_data_sync_mismatch

PiperOrigin-RevId: 929043537
Change-Id: I99c6c67bc02c45dc3a432a585aebac1621757d6d
This commit is contained in:
Copybara-Service
2026-06-09 01:47:11 -07:00
2 changed files with 23 additions and 1 deletions
+4
View File
@@ -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}'
+19 -1
View File
@@ -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('<mujoco></mujoco>')
# 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("""
<mujoco>
<worldbody>
<body>
<freejoint/>
<geom size="0.1"/>
</body>
</worldbody>
</mujoco>
""")
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."""