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:
committed by
Copybara-Service
parent
694eb6bb1d
commit
0e01b51590
@@ -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)
|
||||
|
||||
|
||||
@@ -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."""
|
||||
|
||||
|
||||
Reference in New Issue
Block a user