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)