diff --git a/doc/changelog.rst b/doc/changelog.rst
index 8d911190..22b14cad 100644
--- a/doc/changelog.rst
+++ b/doc/changelog.rst
@@ -8,7 +8,8 @@ Upcoming version (not yet released)
MJX
^^^
-1. Added ``apply_ft``, ``jac``, and ``xfrc_accumulate`` as public functions.
+- Added ``apply_ft``, ``jac``, and ``xfrc_accumulate`` as public functions.
+- Added ``TOUCH`` sensor.
Version 3.2.4 (Oct 15, 2024)
----------------------------
diff --git a/doc/mjx.rst b/doc/mjx.rst
index 9c982edc..c1e0fc16 100644
--- a/doc/mjx.rst
+++ b/doc/mjx.rst
@@ -220,8 +220,8 @@ The following features are **fully supported** in MJX:
- ``MAGNETOMETER``, ``CAMPROJECTION``, ``RANGEFINDER``, ``JOINTPOS``, ``TENDONPOS``, ``ACTUATORPOS``, ``BALLQUAT``,
``FRAMEPOS``, ``FRAMEXAXIS``, ``FRAMEYAXIS``, ``FRAMEZAXIS``, ``FRAMEQUAT``, ``SUBTREECOM``, ``CLOCK``,
``VELOCIMETER``, ``GYRO``, ``JOINTVEL``, ``TENDONVEL``, ``ACTUATORVEL``, ``BALLANGVEL``, ``FRAMELINVEL``,
- ``FRAMEANGVEL``, ``SUBTREELINVEL``, ``SUBTREEANGMOM``, ``ACCELEROMETER``, ``FORCE``, ``TORQUE``, ``ACTUATORFRC``,
- ``JOINTACTFRC``, ``FRAMELINACC``, ``FRAMEANGACC``.
+ ``FRAMEANGVEL``, ``SUBTREELINVEL``, ``SUBTREEANGMOM``, ``TOUCH``, ``ACCELEROMETER``, ``FORCE``, ``TORQUE``,
+ ``ACTUATORFRC``, ``JOINTACTFRC``, ``FRAMELINACC``, ``FRAMEANGACC``.
The following features are **in development** and coming soon:
diff --git a/mjx/mujoco/mjx/_src/ray.py b/mjx/mujoco/mjx/_src/ray.py
index 8e36e1a3..2591385a 100644
--- a/mjx/mujoco/mjx/_src/ray.py
+++ b/mjx/mujoco/mjx/_src/ray.py
@@ -126,6 +126,7 @@ def _ray_box(
p1 = pnt[iface[:, 1]] + x * vec[iface[:, 1]]
valid = jp.abs(p0) <= size[iface[:, 0]]
valid &= jp.abs(p1) <= size[iface[:, 1]]
+ valid &= x >= 0
return jp.min(jp.where(valid, x, jp.inf))
@@ -268,3 +269,20 @@ def ray(
id_ = jp.where(jp.isinf(dists[min_id]), -1, ids[min_id])
return dist, id_
+
+
+def ray_geom(
+ size: jax.Array, pnt: jax.Array, vec: jax.Array, geomtype: GeomType
+) -> jax.Array:
+ """Returns the distance at which a ray intersects with a primitive geom.
+
+ Args:
+ size: geom size (1,), (2,), or (3,)
+ pnt: ray origin point (3,)
+ vec: ray direction (3,)
+ geomtype: type of geom
+
+ Returns:
+ dist: distance from ray origin to geom surface
+ """
+ return _RAY_FUNC[geomtype](size, pnt, vec)
diff --git a/mjx/mujoco/mjx/_src/sensor.py b/mjx/mujoco/mjx/_src/sensor.py
index 02545cc4..e514bdd2 100644
--- a/mjx/mujoco/mjx/_src/sensor.py
+++ b/mjx/mujoco/mjx/_src/sensor.py
@@ -21,6 +21,7 @@ import mujoco
from mujoco.mjx._src import math
from mujoco.mjx._src import ray
from mujoco.mjx._src import smooth
+from mujoco.mjx._src import support
from mujoco.mjx._src.types import Data
from mujoco.mjx._src.types import DisableBit
from mujoco.mjx._src.types import Model
@@ -452,7 +453,65 @@ def sensor_acc(m: Model, d: Data) -> Data:
cutoff = m.sensor_cutoff[idx]
data_type = m.sensor_datatype[idx]
- if sensor_type == SensorType.ACCELEROMETER:
+ if sensor_type == SensorType.TOUCH:
+ # compute contact forces
+ forces = []
+ condim_ids = []
+ for dim in set(d.contact.dim):
+ force, condim_id = support.contact_force_dim(m, d, dim)
+ forces.append(force)
+ condim_ids.append(condim_id)
+ forces = jp.concatenate(forces)[jp.concatenate(condim_ids)]
+
+ # get bodies of contact geoms
+ conbody = jp.array(m.geom_bodyid)[d.contact.geom]
+
+ # get site information
+ site_bodyid = m.site_bodyid[objid]
+ site_size = m.site_size[objid]
+ site_xpos = d.site_xpos[objid]
+ site_xmat = d.site_xmat[objid]
+ site_type = m.site_type[objid]
+ conbody0 = site_bodyid[:, None] == conbody[:, 0]
+ conbody1 = site_bodyid[:, None] == conbody[:, 1]
+ contacts = (d.contact.efc_address >= 0)[None] & (conbody0 | conbody1)
+
+ # compute conray, flip if second body
+ conray = jax.vmap(
+ lambda frame, force: math.normalize(frame[0] * force[0])
+ )(d.contact.frame, forces)
+ conray = jp.where(conbody1[..., None], -conray, conray)
+
+ # compute distance, mapping over sites and contacts
+ def _distance(
+ site_size, site_xpos, site_xmat, site_type, contact_pos, conray
+ ):
+ return jax.vmap(
+ lambda site_size, site_xpos, site_xmat, conray: jax.vmap(
+ lambda pnt, vec: ray.ray_geom(site_size, pnt, vec, site_type)
+ )((contact_pos - site_xpos) @ site_xmat, conray @ site_xmat)
+ )(site_size, site_xpos, site_xmat, conray)
+
+ dist = []
+ dist_id = []
+ for st in set(site_type):
+ (dist_id_site,) = np.nonzero(st == site_type)
+ dist_site = _distance(
+ site_size[dist_id_site],
+ site_xpos[dist_id_site],
+ site_xmat[dist_id_site],
+ st,
+ d.contact.pos,
+ conray[dist_id_site],
+ )
+ dist.append(jp.where(jp.isinf(dist_site), 0, dist_site))
+ dist_id.append(dist_id_site)
+
+ dist = jp.vstack(dist)[np.concatenate(dist_id)]
+
+ # accumulate normal forces for each site
+ sensor = jp.dot((dist > 0) & contacts, forces[:, 0])
+ elif sensor_type == SensorType.ACCELEROMETER:
@jax.vmap
def _accelerometer(cvel, cacc, diff, rot):
diff --git a/mjx/mujoco/mjx/_src/types.py b/mjx/mujoco/mjx/_src/types.py
index 15eaa8c3..f8fdbd23 100644
--- a/mjx/mujoco/mjx/_src/types.py
+++ b/mjx/mujoco/mjx/_src/types.py
@@ -331,6 +331,7 @@ class SensorType(enum.IntEnum):
FRAMEANGVEL: 3D angular velocity
SUBTREELINVEL: subtree linear velocity
SUBTREEANGMOM: subtree angular momentum
+ TOUCH: scalar contact normal forces summed over the sensor zone
ACCELEROMETER: accelerometer
FORCE: force
TORQUE: torque
@@ -363,6 +364,7 @@ class SensorType(enum.IntEnum):
FRAMEANGVEL = mujoco.mjtSensor.mjSENS_FRAMEANGVEL
SUBTREELINVEL = mujoco.mjtSensor.mjSENS_SUBTREELINVEL
SUBTREEANGMOM = mujoco.mjtSensor.mjSENS_SUBTREEANGMOM
+ TOUCH = mujoco.mjtSensor.mjSENS_TOUCH
ACCELEROMETER = mujoco.mjtSensor.mjSENS_ACCELEROMETER
FORCE = mujoco.mjtSensor.mjSENS_FORCE
TORQUE = mujoco.mjtSensor.mjSENS_TORQUE
diff --git a/mjx/mujoco/mjx/test_data/sensor/sensor.xml b/mjx/mujoco/mjx/test_data/sensor/sensor.xml
index e9f1341c..c8257831 100644
--- a/mjx/mujoco/mjx/test_data/sensor/sensor.xml
+++ b/mjx/mujoco/mjx/test_data/sensor/sensor.xml
@@ -27,6 +27,7 @@
-subtreelinvel
-subtreeangmom
* acceleration/force-dependent sensors:
+-touch
-accelerometer
-force
-torque
@@ -100,6 +101,17 @@