diff --git a/mjx/mujoco/mjx/_src/sensor.py b/mjx/mujoco/mjx/_src/sensor.py index 7dabeeb2..e644bcdc 100644 --- a/mjx/mujoco/mjx/_src/sensor.py +++ b/mjx/mujoco/mjx/_src/sensor.py @@ -20,6 +20,7 @@ import mujoco # pylint: disable=g-importing-member from mujoco.mjx._src import math from mujoco.mjx._src.types import Data +from mujoco.mjx._src.types import DisableBit from mujoco.mjx._src.types import Model from mujoco.mjx._src.types import ObjType from mujoco.mjx._src.types import SensorType @@ -30,6 +31,9 @@ import numpy as np def sensor_pos(m: Model, d: Data) -> Data: """Compute position-dependent sensors values.""" + if m.opt.disableflags & DisableBit.SENSOR: + return d + # no position-dependent sensors stage_pos = m.sensor_needstage == mujoco.mjtStage.mjSTAGE_POS if sum(stage_pos) == 0: @@ -145,9 +149,17 @@ def sensor_pos(m: Model, d: Data) -> Data: def sensor_vel(m: Model, d: Data) -> Data: """Compute velocity-dependent sensors values.""" + + if m.opt.disableflags & DisableBit.SENSOR: + return d + return d def sensor_acc(m: Model, d: Data) -> Data: """Compute acceleration/force-dependent sensors values.""" + + if m.opt.disableflags & DisableBit.SENSOR: + return d + return d diff --git a/mjx/mujoco/mjx/_src/sensor_test.py b/mjx/mujoco/mjx/_src/sensor_test.py index d120e757..0ff11a86 100644 --- a/mjx/mujoco/mjx/_src/sensor_test.py +++ b/mjx/mujoco/mjx/_src/sensor_test.py @@ -17,6 +17,7 @@ from absl.testing import absltest from absl.testing import parameterized import jax +from jax import numpy as jp import mujoco from mujoco import mjx from mujoco.mjx._src import test_util @@ -62,6 +63,25 @@ class SensorTest(parameterized.TestCase): # sensor values _assert_eq(d.sensordata, dx.sensordata, 'sensordata') + def test_disable_sensor(self): + """Tests disabling sensor.""" + m = test_util.load_test_file('sensor.xml') + # disable sensors + m.opt.disableflags = m.opt.disableflags | mjx.DisableBit.SENSOR + d = mujoco.MjData(m) + # give the system a little kick to ensure we have non-identity rotations + d.qvel = np.random.random(m.nv) + mujoco.mj_step(m, d, 10) # let dynamics get state significantly non-zero + mx = mjx.put_model(m) + dx = mjx.put_data(m, d) + # random sensor values + random_sensor = jp.array(np.random.random(dx.sensordata.shape)) + dx = dx.replace(sensordata=random_sensor) + # call sensor functions + dx = jax.jit(mjx.forward)(mx, dx) + # sensor values + _assert_eq(random_sensor, dx.sensordata, 'sensordata') + def test_unsupported_sensor(self): """Tests MJX sensor functions do not break for unsupported sensors.""" m = test_util.load_test_file('unsupported_sensor.xml') diff --git a/mjx/mujoco/mjx/_src/types.py b/mjx/mujoco/mjx/_src/types.py index c3f2d2ff..8ac87e41 100644 --- a/mjx/mujoco/mjx/_src/types.py +++ b/mjx/mujoco/mjx/_src/types.py @@ -45,6 +45,7 @@ class DisableBit(enum.IntFlag): WARMSTART: warmstart constraint solver ACTUATION: apply actuation forces REFSAFE: integrator safety: make ref[0]>=2*timestep + SENSOR: sensors """ CONSTRAINT = mujoco.mjtDisableBit.mjDSBL_CONSTRAINT EQUALITY = mujoco.mjtDisableBit.mjDSBL_EQUALITY @@ -56,9 +57,10 @@ class DisableBit(enum.IntFlag): WARMSTART = mujoco.mjtDisableBit.mjDSBL_WARMSTART ACTUATION = mujoco.mjtDisableBit.mjDSBL_ACTUATION REFSAFE = mujoco.mjtDisableBit.mjDSBL_REFSAFE + SENSOR = mujoco.mjtDisableBit.mjDSBL_SENSOR EULERDAMP = mujoco.mjtDisableBit.mjDSBL_EULERDAMP FILTERPARENT = mujoco.mjtDisableBit.mjDSBL_FILTERPARENT - # unsupported: FRICTIONLOSS, SENSOR, MIDPHASE + # unsupported: FRICTIONLOSS, MIDPHASE class JointType(enum.IntEnum):