diff --git a/mjx/mujoco/mjx/_src/io.py b/mjx/mujoco/mjx/_src/io.py index 6acef205..cb9d395a 100644 --- a/mjx/mujoco/mjx/_src/io.py +++ b/mjx/mujoco/mjx/_src/io.py @@ -514,7 +514,13 @@ def get_data_into( if restricted_to in ('mujoco', 'mjx'): continue # don't copy fields that are mujoco-only or MJX-only else: - getattr(result_i, field.name)[:] = value + result_field = getattr(result_i, field.name) + if result_field.shape != value.shape: + raise ValueError( + f'Input field {field.name} has shape {value.shape}, but output' + f' has shape {result_field.shape}' + ) + result_field[:] = value else: setattr(result_i, field.name, value) diff --git a/mjx/mujoco/mjx/_src/io_test.py b/mjx/mujoco/mjx/_src/io_test.py index 01977792..494b2753 100644 --- a/mjx/mujoco/mjx/_src/io_test.py +++ b/mjx/mujoco/mjx/_src/io_test.py @@ -482,6 +482,18 @@ class DataIOTest(parameterized.TestCase): self.assertEqual(d_2.contact.frame.shape, (1, 9)) np.testing.assert_allclose(d_2.contact.frame, d.contact.frame) + def test_get_data_into_wrong_shape(self): + """Tests that get_data_into throwsif input and output shapes don't match.""" + + m = mujoco.MjModel.from_xml_string(_MULTIPLE_CONSTRAINTS) + d = mujoco.MjData(m) + mujoco.mj_step(m, d, 2) + dx = mjx.put_data(m, d) + m_2 = mujoco.MjModel.from_xml_string(_MULTIPLE_CONVEX_OBJECTS) + d_2 = mujoco.MjData(m_2) + with self.assertRaisesRegex(ValueError, r'Input field.*has shape.*'): + mjx.get_data_into(d_2, m, dx) + def test_make_matches_put(self): """Test that make_data produces a pytree that matches put_data."""