diff --git a/mjx/mujoco/mjx/_src/support.py b/mjx/mujoco/mjx/_src/support.py index dead920c..c5ae3157 100644 --- a/mjx/mujoco/mjx/_src/support.py +++ b/mjx/mujoco/mjx/_src/support.py @@ -437,6 +437,10 @@ class BindData(object): def set(self, name: str, value: jax.Array) -> Data: """Set the value of an array in an MJX Data.""" array = getattr(self.data, self.__getname(name)) + try: + iter(value) + except TypeError: + value = [value] if len(value) == 1: array = array.at[self.id].set(value[0]) else: diff --git a/mjx/mujoco/mjx/_src/support_test.py b/mjx/mujoco/mjx/_src/support_test.py index 567d2cc7..53c69cf9 100644 --- a/mjx/mujoco/mjx/_src/support_test.py +++ b/mjx/mujoco/mjx/_src/support_test.py @@ -236,6 +236,9 @@ class SupportTest(parameterized.TestCase): dx4 = dx.bind(mx, s.actuators[1]).set('ctrl', [6]) np.testing.assert_array_equal(dx4.bind(mx, s.actuators).ctrl, [0, 6, 0]) np.testing.assert_array_equal(dx.bind(mx, s.actuators).ctrl, [0, 0, 0]) + dx5 = dx.bind(mx, s.actuators[1]).set('ctrl', 7) + np.testing.assert_array_equal(dx5.bind(mx, s.actuators).ctrl, [0, 7, 0]) + np.testing.assert_array_equal(dx.bind(mx, s.actuators).ctrl, [0, 0, 0]) # test invalid name with self.assertRaises(AttributeError):