Add support for disabling MJX sensors.
PiperOrigin-RevId: 664410169 Change-Id: Icb278e9d7a9a493d5a61ab04c1300d8e412d0e83
This commit is contained in:
committed by
Copybara-Service
parent
0bffd744f9
commit
6acf40613d
@@ -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
|
||||
|
||||
@@ -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')
|
||||
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user