diff --git a/mjx/mujoco/mjx/_src/scan.py b/mjx/mujoco/mjx/_src/scan.py index 496ba4cd..0bcfc195 100644 --- a/mjx/mujoco/mjx/_src/scan.py +++ b/mjx/mujoco/mjx/_src/scan.py @@ -132,6 +132,8 @@ def _nvmap(f: Callable[..., Y], *args) -> Y: def _check_input(m: Model, args: Any, in_types: str) -> None: """Checks that scan input has the right shape.""" + if m.nv == 0: + raise ValueError('Scan across Model with zero DoFs unsupported.') size = { 'b': m.nbody, 'j': m.njnt, diff --git a/mjx/mujoco/mjx/_src/scan_test.py b/mjx/mujoco/mjx/_src/scan_test.py index 74bf3224..758ebe89 100644 --- a/mjx/mujoco/mjx/_src/scan_test.py +++ b/mjx/mujoco/mjx/_src/scan_test.py @@ -61,10 +61,7 @@ class ScanTest(absltest.TestCase): return body_id + 1 b_in = jp.array([1]) - b_expect = jp.array([2]) - b_out = scan.flat(m, fn, 'b', 'b', b_in) - - np.testing.assert_equal(np.array(b_out), np.array(b_expect)) + self.assertRaises(ValueError, scan.flat, m, fn, 'b', 'b', b_in) def test_flat_joints(self): """Tests scanning over bodies with joints of different types."""