From da32b0db30476c43897b6f38b9f64dc05f863114 Mon Sep 17 00:00:00 2001 From: Alessio Quaglino Date: Fri, 10 Jan 2025 08:07:05 -0800 Subject: [PATCH] Handle `ctrl` special case explicitly in bind(). PiperOrigin-RevId: 714055763 Change-Id: I020008723af4a590926268a2ea615af54e20aa0b --- mjx/mujoco/mjx/_src/support.py | 19 ++++++++++--------- mjx/mujoco/mjx/_src/support_test.py | 6 ++++-- 2 files changed, 14 insertions(+), 11 deletions(-) diff --git a/mjx/mujoco/mjx/_src/support.py b/mjx/mujoco/mjx/_src/support.py index 18992535..9032b8d2 100644 --- a/mjx/mujoco/mjx/_src/support.py +++ b/mjx/mujoco/mjx/_src/support.py @@ -404,7 +404,7 @@ class BindData(object): self.prefix = 'cam_' ids.append(name2id(model, mujoco.mjtObj.mjOBJ_CAMERA, spec.name)) case mujoco.MjsTendon(): - self.prefix = 'tendon_' + self.prefix = 'ten_' ids.append(name2id(model, mujoco.mjtObj.mjOBJ_TENDON, spec.name)) case mujoco.MjsActuator(): self.prefix = 'actuator_' @@ -412,6 +412,9 @@ class BindData(object): case mujoco.MjsSensor(): self.prefix = 'sensor_' ids.append(name2id(model, mujoco.mjtObj.mjOBJ_SENSOR, spec.name)) + case mujoco.MjsEquality(): + self.prefix = 'eq_' + ids.append(name2id(model, mujoco.mjtObj.mjOBJ_EQUALITY, spec.name)) case _: raise ValueError('invalid spec type') if len(ids) == 1: @@ -420,15 +423,13 @@ class BindData(object): self.id = ids def __getname(self, name: str): - try: - getattr(self.data, self.prefix + name) - return self.prefix + name - except AttributeError: - try: - getattr(self.data, name) + if name == 'ctrl': + if self.prefix == 'actuator_': return name - except AttributeError as e: - raise ValueError(f'invalid name: {name}') from e + else: + raise AttributeError('ctrl is not available for this type') + else: + return self.prefix + name def __getattr__(self, name: str): return getattr(self.data, self.__getname(name))[self.id, ...] diff --git a/mjx/mujoco/mjx/_src/support_test.py b/mjx/mujoco/mjx/_src/support_test.py index 9add48b3..567d2cc7 100644 --- a/mjx/mujoco/mjx/_src/support_test.py +++ b/mjx/mujoco/mjx/_src/support_test.py @@ -238,9 +238,11 @@ class SupportTest(parameterized.TestCase): np.testing.assert_array_equal(dx.bind(mx, s.actuators).ctrl, [0, 0, 0]) # test invalid name - with self.assertRaises(ValueError): + with self.assertRaises(AttributeError): + print(dx.bind(mx, s.geoms).ctrl) + with self.assertRaises(AttributeError): print(dx.bind(mx, s.actuators).actuator_ctrl) - with self.assertRaises(ValueError): + with self.assertRaises(AttributeError): print(dx.bind(mx, s.actuators).set('actuator_ctrl', [1, 2, 3])) _CONTACTS = """