diff --git a/mjx/mujoco/mjx/_src/support.py b/mjx/mujoco/mjx/_src/support.py index 10ec9d2a..905a2cbf 100644 --- a/mjx/mujoco/mjx/_src/support.py +++ b/mjx/mujoco/mjx/_src/support.py @@ -381,6 +381,7 @@ class BindData(object): def __init__(self, data: Data, model: Model, specs: Sequence[Any]): self.data = data + self.model = model try: iter(specs) except TypeError: @@ -438,10 +439,22 @@ class BindData(object): return self.prefix + name def __getattr__(self, name: str): + if name == 'sensordata': + adr = self.model.sensor_adr[self.id] + num = self.model.sensor_dim[self.id] + if isinstance(self.id, list): + idx = [] + for i, n in zip(self.id, num): + idx.extend(adr[i] + j for j in range(n)) + return getattr(self.data, name)[idx, ...] + else: + return getattr(self.data, name)[adr : adr + num, ...] return getattr(self.data, self.__getname(name))[self.id, ...] def set(self, name: str, value: jax.Array) -> Data: """Set the value of an array in an MJX Data.""" + if name == 'sensordata': + raise AttributeError('sensordata is readonly') array = getattr(self.data, self.__getname(name)) try: iter(value) diff --git a/mjx/mujoco/mjx/_src/support_test.py b/mjx/mujoco/mjx/_src/support_test.py index 5cfd8cd0..a61374da 100644 --- a/mjx/mujoco/mjx/_src/support_test.py +++ b/mjx/mujoco/mjx/_src/support_test.py @@ -180,6 +180,12 @@ class SupportTest(parameterized.TestCase): + + + + + + """ @@ -224,6 +230,15 @@ class SupportTest(parameterized.TestCase): dx.bind(mx, s.actuators[i]).ctrl, d.ctrl[i] ) + np.testing.assert_array_equal( + dx.bind(mx, s.sensors).sensordata, d.sensordata + ) + for i in range(m.nsensor): + np.testing.assert_array_equal( + dx.bind(mx, s.sensors[i]).sensordata, + d.sensordata[m.sensor_adr[i] : m.sensor_adr[i] + m.sensor_dim[i]], + ) + # test setting np.testing.assert_array_equal(d.ctrl, [0, 0, 0]) np.testing.assert_array_equal(dx.bind(mx, s.actuators).ctrl, d.ctrl)