From f1d557c12517b0e0af7eef8d4e45dffeb12591a7 Mon Sep 17 00:00:00 2001 From: Alessio Quaglino Date: Wed, 15 Jan 2025 09:36:44 -0800 Subject: [PATCH] Support for MJX bind set() for scalars. PiperOrigin-RevId: 715832281 Change-Id: I217021e81fae95e44e9297a1c447891ff76572bd --- mjx/mujoco/mjx/_src/support.py | 4 ++++ mjx/mujoco/mjx/_src/support_test.py | 3 +++ 2 files changed, 7 insertions(+) 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):