From fa156ed22f4baa701654b3dc588b09726694885b Mon Sep 17 00:00:00 2001 From: Taylor Howell Date: Thu, 5 Sep 2024 06:55:51 -0700 Subject: [PATCH] Add accelerometer to MJX. PiperOrigin-RevId: 671355424 Change-Id: I8593bcad6ba311f3c555b118889e2a5c6b367ec7 --- doc/changelog.rst | 2 +- mjx/mujoco/mjx/_src/sensor.py | 27 ++++++++++++++++++++-- mjx/mujoco/mjx/_src/sensor_test.py | 22 +++++++++++++----- mjx/mujoco/mjx/_src/types.py | 2 ++ mjx/mujoco/mjx/test_data/sensor/sensor.xml | 3 +++ 5 files changed, 47 insertions(+), 9 deletions(-) diff --git a/doc/changelog.rst b/doc/changelog.rst index d045ca61..5d24e8e4 100644 --- a/doc/changelog.rst +++ b/doc/changelog.rst @@ -54,7 +54,7 @@ MJX ``SUBTREECOM``, ``CLOCK``. - Added velocity-dependent sensors: ``VELOCIMETER``, ``GYRO``, ``JOINTVEL``, ``ACTUATORVEL``, ``BALLANGVEL``, ``FRAMELINVEL``, ``FRAMEANGVEL``, ``SUBTREELINVEL``, ``SUBTREEANGMOM``. -- Added acceleration/force-dependent sensors: ``ACTUATORFRC``, ``JOINTACTFRC``. +- Added acceleration/force-dependent sensors: ``ACCELEROMETER``, ``ACTUATORFRC``, ``JOINTACTFRC``. - Changed default policy to avoid placing unused (MuJoCo-only) arrays on device. - Added ``device`` parameter to ``mjx.make_data`` to bring it to parity with ``mjx.put_model`` and ``mjx.put_data``. - Added support for :ref:`implicitfast integration` for all cases except diff --git a/mjx/mujoco/mjx/_src/sensor.py b/mjx/mujoco/mjx/_src/sensor.py index 7eb1b315..009d7d7a 100644 --- a/mjx/mujoco/mjx/_src/sensor.py +++ b/mjx/mujoco/mjx/_src/sensor.py @@ -420,16 +420,39 @@ def sensor_acc(m: Model, d: Data) -> Data: return d stage_acc = m.sensor_needstage == mujoco.mjtStage.mjSTAGE_ACC + sensor_types = set(m.sensor_type[stage_acc]) + + if sensor_types & {SensorType.ACCELEROMETER}: + d = smooth.rne_postconstraint(m, d) + sensors, adrs = [], [] - for sensor_type in set(m.sensor_type[stage_acc]): + for sensor_type in sensor_types: idx = m.sensor_type == sensor_type objid = m.sensor_objid[idx] adr = m.sensor_adr[idx] cutoff = m.sensor_cutoff[idx] data_type = m.sensor_datatype[idx] - if sensor_type == SensorType.ACTUATORFRC: + if sensor_type == SensorType.ACCELEROMETER: + + @jax.vmap + def _accelerometer(cvel, cacc, diff, rot): + ang = rot.T @ cvel[:3] + lin = rot.T @ (cvel[3:] - jp.cross(diff, cvel[:3])) + acc = rot.T @ (cacc[3:] - jp.cross(diff, cacc[:3])) + correction = jp.cross(ang, lin) + return acc + correction + + bodyid = m.site_bodyid[objid] + rot = d.site_xmat[objid] + cvel = d.cvel[bodyid] + cacc = d.cacc[bodyid] + dif = d.site_xpos[objid] - d.subtree_com[m.body_rootid[bodyid]] + + sensor = _accelerometer(cvel, cacc, dif, rot) + adr = (adr[:, None] + np.arange(3)[None]).reshape(-1) + elif sensor_type == SensorType.ACTUATORFRC: sensor = d.actuator_force[objid] elif sensor_type == SensorType.JOINTACTFRC: sensor = d.qfrc_actuator[m.jnt_dofadr[objid]] diff --git a/mjx/mujoco/mjx/_src/sensor_test.py b/mjx/mujoco/mjx/_src/sensor_test.py index e2b7bd35..42c2352a 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 @@ -46,18 +47,27 @@ class SensorTest(parameterized.TestCase): m = test_util.load_test_file(filename) d = mujoco.MjData(m) # give the system a little kick to ensure we have non-identity rotations - d.qvel = np.random.random(m.nv) + d.qvel = 0.1 * np.random.random(m.nv) + # apply external forces + d.xfrc_applied = 0.1 * np.random.random(d.xfrc_applied.shape) # apply control for activation dynamics d.ctrl = np.clip( - np.random.random(m.nu), + 0.1 * np.random.random(m.nu), m.actuator_ctrlrange[:, 0], m.actuator_ctrlrange[:, 1], ) - mujoco.mj_step(m, d, 10) # let dynamics get state significantly non-zero - mx = mjx.put_model(m) - dx = mjx.put_data(m, d) - + mujoco.mj_step(m, d, 100) mujoco.mj_forward(m, d) + + mx = mjx.put_model(m) + dx = mjx.put_data(m, d).replace( + sensordata=jp.zeros_like(d.sensordata), + subtree_linvel=jp.zeros_like(d.subtree_linvel), + subtree_angmom=jp.zeros_like(d.subtree_angmom), + cacc=jp.zeros_like(d.cacc), + cfrc_int=jp.zeros_like(d.cfrc_int), + cfrc_ext=jp.zeros_like(d.cfrc_ext), + ) dx = jax.jit(mjx.forward)(mx, dx) # sensor values diff --git a/mjx/mujoco/mjx/_src/types.py b/mjx/mujoco/mjx/_src/types.py index def0899b..6f1da229 100644 --- a/mjx/mujoco/mjx/_src/types.py +++ b/mjx/mujoco/mjx/_src/types.py @@ -318,6 +318,7 @@ class SensorType(enum.IntEnum): FRAMEANGVEL: 3D angular velocity SUBTREELINVEL: subtree linear velocity SUBTREEANGMOM: subtree angular momentum + ACCELEROMETER: accelerometer ACTUATORFRC: scalar actuator force JOINTACTFRC: scalar actuator force, measured at the joint """ @@ -343,6 +344,7 @@ class SensorType(enum.IntEnum): FRAMEANGVEL = mujoco.mjtSensor.mjSENS_FRAMEANGVEL SUBTREELINVEL = mujoco.mjtSensor.mjSENS_SUBTREELINVEL SUBTREEANGMOM = mujoco.mjtSensor.mjSENS_SUBTREEANGMOM + ACCELEROMETER = mujoco.mjtSensor.mjSENS_ACCELEROMETER ACTUATORFRC = mujoco.mjtSensor.mjSENS_ACTUATORFRC JOINTACTFRC = mujoco.mjtSensor.mjSENS_JOINTACTFRC diff --git a/mjx/mujoco/mjx/test_data/sensor/sensor.xml b/mjx/mujoco/mjx/test_data/sensor/sensor.xml index 0c5ca76b..5449f091 100644 --- a/mjx/mujoco/mjx/test_data/sensor/sensor.xml +++ b/mjx/mujoco/mjx/test_data/sensor/sensor.xml @@ -25,6 +25,7 @@ -subtreelinvel -subtreeangmom * acceleration/force-dependent sensors: +-accelerometer -actuatorfrc -jointactfrc --> @@ -93,6 +94,7 @@ + @@ -140,6 +142,7 @@ +