From 9d732117575ca9372ef31fe065669e1e4c2d8604 Mon Sep 17 00:00:00 2001 From: Taylor Howell Date: Mon, 16 Sep 2024 06:56:08 -0700 Subject: [PATCH] Add force and torque sensors to MJX. PiperOrigin-RevId: 675130978 Change-Id: I8a3b925c57f317aebc0c345578cebeb4210dc11d --- doc/changelog.rst | 3 ++- doc/mjx.rst | 2 +- mjx/mujoco/mjx/_src/sensor.py | 24 +++++++++++++++++++++- mjx/mujoco/mjx/_src/sensor_test.py | 5 +++-- mjx/mujoco/mjx/_src/types.py | 4 ++++ mjx/mujoco/mjx/test_data/sensor/model.xml | 16 +++++++++++++++ mjx/mujoco/mjx/test_data/sensor/sensor.xml | 20 ++++++++++++++++++ 7 files changed, 69 insertions(+), 5 deletions(-) diff --git a/doc/changelog.rst b/doc/changelog.rst index 89e95d49..149571fe 100644 --- a/doc/changelog.rst +++ b/doc/changelog.rst @@ -63,7 +63,8 @@ MJX ``SUBTREECOM``, ``CLOCK``. 14. Added velocity-dependent sensors: ``VELOCIMETER``, ``GYRO``, ``JOINTVEL``, ``ACTUATORVEL``, ``BALLANGVEL``, ``FRAMELINVEL``, ``FRAMEANGVEL``, ``SUBTREELINVEL``, ``SUBTREEANGMOM``. -15. Added acceleration/force-dependent sensors: ``ACCELEROMETER``, ``ACTUATORFRC``, ``JOINTACTFRC``. +15. Added acceleration/force-dependent sensors: ``ACCELEROMETER``, ``FORCE``, ``TORQUE``, ``ACTUATORFRC``, + ``JOINTACTFRC``. 16. Changed default policy to avoid placing unused (MuJoCo-only) arrays on device. 17. Added ``device`` parameter to ``mjx.make_data`` to bring it to parity with ``mjx.put_model`` and ``mjx.put_data``. 18. Added support for :ref:`implicitfast integration` for all cases except diff --git a/doc/mjx.rst b/doc/mjx.rst index 1a82c892..abfdb690 100644 --- a/doc/mjx.rst +++ b/doc/mjx.rst @@ -220,7 +220,7 @@ The following features are **fully supported** in MJX: - ``MAGNETOMETER``, ``CAMPROJECTION``, ``RANGEFINDER``, ``JOINTPOS``, ``ACTUATORPOS``, ``BALLQUAT``, ``FRAMEPOS``, ``FRAMEXAXIS``, ``FRAMEYAXIS``, ``FRAMEZAXIS``, ``FRAMEQUAT``, ``SUBTREECOM``, ``CLOCK``, ``VELOCIMETER``, ``GYRO``, ``JOINTVEL``, ``ACTUATORVEL``, ``BALLANGVEL``, ``FRAMELINVEL``, ``FRAMEANGVEL``, ``SUBTREELINVEL``, - ``SUBTREEANGMOM``, ``ACCELEROMETER``, ``ACTUATORFRC``, ``JOINTACTFRC``. + ``SUBTREEANGMOM``, ``ACCELEROMETER``, ``FORCE``, ``TORQUE``, ``ACTUATORFRC``, ``JOINTACTFRC``. The following features are **in development** and coming soon: diff --git a/mjx/mujoco/mjx/_src/sensor.py b/mjx/mujoco/mjx/_src/sensor.py index 009d7d7a..07811d3c 100644 --- a/mjx/mujoco/mjx/_src/sensor.py +++ b/mjx/mujoco/mjx/_src/sensor.py @@ -422,7 +422,11 @@ def sensor_acc(m: Model, d: Data) -> Data: stage_acc = m.sensor_needstage == mujoco.mjtStage.mjSTAGE_ACC sensor_types = set(m.sensor_type[stage_acc]) - if sensor_types & {SensorType.ACCELEROMETER}: + if sensor_types & { + SensorType.ACCELEROMETER, + SensorType.FORCE, + SensorType.TORQUE, + }: d = smooth.rne_postconstraint(m, d) sensors, adrs = [], [] @@ -452,6 +456,24 @@ def sensor_acc(m: Model, d: Data) -> Data: sensor = _accelerometer(cvel, cacc, dif, rot) adr = (adr[:, None] + np.arange(3)[None]).reshape(-1) + elif sensor_type == SensorType.FORCE: + bodyid = m.site_bodyid[objid] + cfrc_int = d.cfrc_int[bodyid] + site_xmat = d.site_xmat[objid] + sensor = jax.vmap(lambda mat, vec: mat.T @ vec)( + site_xmat, cfrc_int[:, 3:] + ) + adr = (adr[:, None] + np.arange(3)[None]).reshape(-1) + elif sensor_type == SensorType.TORQUE: + bodyid = m.site_bodyid[objid] + rootid = m.body_rootid[bodyid] + cfrc_int = d.cfrc_int[bodyid] + site_xmat = d.site_xmat[objid] + dif = d.site_xpos[objid] - d.subtree_com[rootid] + sensor = jax.vmap( + lambda vec, dif, rot: rot.T @ (vec[:3] - jp.cross(dif, vec[3:])) + )(cfrc_int, dif, site_xmat) + adr = (adr[:, None] + np.arange(3)[None]).reshape(-1) elif sensor_type == SensorType.ACTUATORFRC: sensor = d.actuator_force[objid] elif sensor_type == SensorType.JOINTACTFRC: diff --git a/mjx/mujoco/mjx/_src/sensor_test.py b/mjx/mujoco/mjx/_src/sensor_test.py index 4613b527..9eedac13 100644 --- a/mjx/mujoco/mjx/_src/sensor_test.py +++ b/mjx/mujoco/mjx/_src/sensor_test.py @@ -68,9 +68,10 @@ class SensorTest(parameterized.TestCase): cfrc_int=jp.zeros_like(d.cfrc_int), cfrc_ext=jp.zeros_like(d.cfrc_ext), ) - dx = jax.jit(mjx.forward)(mx, dx) + dx = jax.jit(mjx.sensor_pos)(mx, dx) + dx = jax.jit(mjx.sensor_vel)(mx, dx) + dx = jax.jit(mjx.sensor_acc)(mx, dx) - # sensor values _assert_eq(d.sensordata, dx.sensordata, 'sensordata') def test_disable_sensor(self): diff --git a/mjx/mujoco/mjx/_src/types.py b/mjx/mujoco/mjx/_src/types.py index 6a200e1c..17f5bf2e 100644 --- a/mjx/mujoco/mjx/_src/types.py +++ b/mjx/mujoco/mjx/_src/types.py @@ -321,6 +321,8 @@ class SensorType(enum.IntEnum): SUBTREELINVEL: subtree linear velocity SUBTREEANGMOM: subtree angular momentum ACCELEROMETER: accelerometer + FORCE: force + TORQUE: torque ACTUATORFRC: scalar actuator force JOINTACTFRC: scalar actuator force, measured at the joint """ @@ -347,6 +349,8 @@ class SensorType(enum.IntEnum): SUBTREELINVEL = mujoco.mjtSensor.mjSENS_SUBTREELINVEL SUBTREEANGMOM = mujoco.mjtSensor.mjSENS_SUBTREEANGMOM ACCELEROMETER = mujoco.mjtSensor.mjSENS_ACCELEROMETER + FORCE = mujoco.mjtSensor.mjSENS_FORCE + TORQUE = mujoco.mjtSensor.mjSENS_TORQUE ACTUATORFRC = mujoco.mjtSensor.mjSENS_ACTUATORFRC JOINTACTFRC = mujoco.mjtSensor.mjSENS_JOINTACTFRC diff --git a/mjx/mujoco/mjx/test_data/sensor/model.xml b/mjx/mujoco/mjx/test_data/sensor/model.xml index 18726d46..f152f99b 100644 --- a/mjx/mujoco/mjx/test_data/sensor/model.xml +++ b/mjx/mujoco/mjx/test_data/sensor/model.xml @@ -47,6 +47,22 @@ + + + + + + + + + + + + + + + + diff --git a/mjx/mujoco/mjx/test_data/sensor/sensor.xml b/mjx/mujoco/mjx/test_data/sensor/sensor.xml index 5449f091..79129f2f 100644 --- a/mjx/mujoco/mjx/test_data/sensor/sensor.xml +++ b/mjx/mujoco/mjx/test_data/sensor/sensor.xml @@ -26,6 +26,8 @@ -subtreeangmom * acceleration/force-dependent sensors: -accelerometer +-force +-torque -actuatorfrc -jointactfrc --> @@ -78,6 +80,22 @@ + + + + + + + + + + + + + + + + @@ -88,6 +106,7 @@ + @@ -95,6 +114,7 @@ +