Output a useful message when there's a mismatch between MJX data and MuJoCo data array sizes.

PiperOrigin-RevId: 725151186
Change-Id: I71f81bb82e2b8667880a31db5913e98f7885e677
This commit is contained in:
Nimrod Gileadi
2025-02-10 03:50:47 -08:00
committed by Copybara-Service
parent 694eb6bb1d
commit 0e01b51590
2 changed files with 19 additions and 1 deletions
+7 -1
View File
@@ -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)
+12
View File
@@ -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."""