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 @@ + + + + + + + + + + + @@ -121,6 +133,7 @@ + @@ -128,6 +141,7 @@ + @@ -138,6 +152,7 @@ +