From 51733c2a8ac081ed7c44d8e96bb0f8d6d30059f0 Mon Sep 17 00:00:00 2001 From: Taylor Howell Date: Sun, 18 Aug 2024 05:15:10 -0700 Subject: [PATCH] Add rangefinder sensor to MJX. PiperOrigin-RevId: 664417458 Change-Id: I1ad7720e48541349f155984584314b2b49ccbf36 --- doc/changelog.rst | 4 ++-- mjx/mujoco/mjx/_src/sensor.py | 14 ++++++++++++++ mjx/mujoco/mjx/_src/types.py | 2 ++ mjx/mujoco/mjx/test_data/sensor.xml | 15 ++++++++++++++- 4 files changed, 32 insertions(+), 3 deletions(-) diff --git a/doc/changelog.rst b/doc/changelog.rst index 1fd5ebf1..9152b8a7 100644 --- a/doc/changelog.rst +++ b/doc/changelog.rst @@ -17,8 +17,8 @@ General MJX ^^^ 5. Added ``efc_pos`` to ``mjx.Data`` (:github:issue:`1388`). -6. Added position-dependent sensors: ``MAGNETOMETER``, ``JOINTPOS``, ``ACTUATORPOS``, ``BALLQUAT``, ``FRAMEPOS``, - ``FRAMEXAXIS``, ``FRAMEYAXIS``, ``FRAMEZAXIS``, ``SUBTREECOM``, ``CLOCK``. +6. Added position-dependent sensors: ``MAGNETOMETER``, ``RANGEFINDER``, ``JOINTPOS``, ``ACTUATORPOS``, ``BALLQUAT``, + ``FRAMEPOS``, ``FRAMEXAXIS``, ``FRAMEYAXIS``, ``FRAMEZAXIS``, ``SUBTREECOM``, ``CLOCK``. 7. Changed default policy to avoid placing unused (MuJoCo-only) arrays on device. 8. Added ``device`` parameter to ``mjx.make_data`` to bring it to parity with ``mjx.put_model`` and ``mjx.put_data``. diff --git a/mjx/mujoco/mjx/_src/sensor.py b/mjx/mujoco/mjx/_src/sensor.py index e644bcdc..afe3741d 100644 --- a/mjx/mujoco/mjx/_src/sensor.py +++ b/mjx/mujoco/mjx/_src/sensor.py @@ -19,6 +19,7 @@ from jax import numpy as jp import mujoco # pylint: disable=g-importing-member from mujoco.mjx._src import math +from mujoco.mjx._src import ray from mujoco.mjx._src.types import Data from mujoco.mjx._src.types import DisableBit from mujoco.mjx._src.types import Model @@ -71,6 +72,19 @@ def sensor_pos(m: Model, d: Data) -> Data: d.site_xmat[objid] ).reshape(-1) adr = (adr[:, None] + np.arange(3)[None]).reshape(-1) + elif sensor_type == SensorType.RANGEFINDER: + site_bodyid = m.site_bodyid[objid] + for sid in set(site_bodyid): + id_ = sid == site_bodyid + objid_ = objid[id_] + site_xpos = d.site_xpos[objid_] + site_mat = d.site_xmat[objid_].reshape((-1, 9))[:, np.array([2, 5, 8])] + sensor, _ = jax.vmap( + ray.ray, in_axes=(None, None, 0, 0, None, None, None) + )(m, d, site_xpos, site_mat, (), True, sid) + sensors.append(sensor) + adrs.append(adr[id_]) + continue # avoid adding to sensors/adrs list a second time elif sensor_type == SensorType.JOINTPOS: sensor = d.qpos[m.jnt_qposadr[objid]] elif sensor_type == SensorType.ACTUATORPOS: diff --git a/mjx/mujoco/mjx/_src/types.py b/mjx/mujoco/mjx/_src/types.py index 8ac87e41..b92db752 100644 --- a/mjx/mujoco/mjx/_src/types.py +++ b/mjx/mujoco/mjx/_src/types.py @@ -293,6 +293,7 @@ class SensorType(enum.IntEnum): Members: MAGNETOMETER: magnetometer + RANGEFINDER: rangefinder JOINTPOS: joint position ACTUATORPOS: actuator position BALLQUAT: ball joint orientation @@ -304,6 +305,7 @@ class SensorType(enum.IntEnum): CLOCK: simulation time """ MAGNETOMETER = mujoco.mjtSensor.mjSENS_MAGNETOMETER + RANGEFINDER = mujoco.mjtSensor.mjSENS_RANGEFINDER JOINTPOS = mujoco.mjtSensor.mjSENS_JOINTPOS ACTUATORPOS = mujoco.mjtSensor.mjSENS_ACTUATORPOS BALLQUAT = mujoco.mjtSensor.mjSENS_BALLQUAT diff --git a/mjx/mujoco/mjx/test_data/sensor.xml b/mjx/mujoco/mjx/test_data/sensor.xml index 13178604..80763173 100644 --- a/mjx/mujoco/mjx/test_data/sensor.xml +++ b/mjx/mujoco/mjx/test_data/sensor.xml @@ -2,6 +2,7 @@ * position-dependent sensors: -magnetometer +-rangefinder -jointpos -actuatorpos -ballquat @@ -15,11 +16,16 @@ * acceleration/force-dependent sensors: --> + + + - + + + @@ -39,6 +45,11 @@ + + + + + @@ -51,6 +62,7 @@ + @@ -67,6 +79,7 @@ +