From 0e01b515902f18410cde3b8f02c8892b4b916e95 Mon Sep 17 00:00:00 2001 From: Nimrod Gileadi Date: Mon, 10 Feb 2025 03:50:47 -0800 Subject: [PATCH] Output a useful message when there's a mismatch between MJX data and MuJoCo data array sizes. PiperOrigin-RevId: 725151186 Change-Id: I71f81bb82e2b8667880a31db5913e98f7885e677 --- mjx/mujoco/mjx/_src/io.py | 8 +++++++- mjx/mujoco/mjx/_src/io_test.py | 12 ++++++++++++ 2 files changed, 19 insertions(+), 1 deletion(-) 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."""