Add support for disabling MJX sensors.

PiperOrigin-RevId: 664410169
Change-Id: Icb278e9d7a9a493d5a61ab04c1300d8e412d0e83
This commit is contained in:
Taylor Howell
2024-08-18 04:40:24 -07:00
committed by Copybara-Service
parent 0bffd744f9
commit 6acf40613d
3 changed files with 35 additions and 1 deletions
+12
View File
@@ -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
+20
View File
@@ -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')
+3 -1
View File
@@ -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):