diff --git a/doc/mjx.rst b/doc/mjx.rst
index 23dd1864..499f9056 100644
--- a/doc/mjx.rst
+++ b/doc/mjx.rst
@@ -243,8 +243,9 @@ 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``, ``TOUCH``, ``ACCELEROMETER``, ``FORCE``, ``TORQUE``,
- ``ACTUATORFRC``, ``JOINTACTFRC``, ``TENDONACTFRC``, ``FRAMELINACC``, ``FRAMEANGACC``
+ ``FRAMEANGVEL``, ``SUBTREELINVEL``, ``SUBTREEANGMOM``, ``TOUCH``, ``CONTACT``, ``ACCELEROMETER``, ``FORCE``,
+ ``TORQUE``, ``ACTUATORFRC``, ``JOINTACTFRC``, ``TENDONACTFRC``, ``FRAMELINACC``, ``FRAMEANGACC``
+ (``CONTACT``: matching ``none-none``, ``geom-geom``; reduction ``mindist``, ``maxforce``; data ``all``)
* - Lights
- Positions and directions of lights
diff --git a/mjx/mujoco/mjx/_src/forward.py b/mjx/mujoco/mjx/_src/forward.py
index d8faf293..c760b0ff 100644
--- a/mjx/mujoco/mjx/_src/forward.py
+++ b/mjx/mujoco/mjx/_src/forward.py
@@ -432,13 +432,13 @@ def forward(m: Model, d: Data) -> Data:
d = sensor.sensor_vel(m, d)
d = fwd_actuation(m, d)
d = fwd_acceleration(m, d)
- d = sensor.sensor_acc(m, d)
if d._impl.efc_J.size == 0:
d = d.replace(qacc=d.qacc_smooth)
return d
d = named_scope(solver.solve)(m, d)
+ d = sensor.sensor_acc(m, d)
return d
diff --git a/mjx/mujoco/mjx/_src/io.py b/mjx/mujoco/mjx/_src/io.py
index bf64260c..c690011c 100644
--- a/mjx/mujoco/mjx/_src/io.py
+++ b/mjx/mujoco/mjx/_src/io.py
@@ -230,6 +230,41 @@ def _put_model_jax(
if m.nflex:
raise NotImplementedError('Flex not implemented for JAX backend.')
+ # contact sensor
+ is_contact_sensor = m.sensor_type == types.SensorType.CONTACT
+ if is_contact_sensor.any():
+ objtype = m.sensor_objtype[is_contact_sensor]
+ reftype = m.sensor_reftype[is_contact_sensor]
+ contact_sensor_type = set(np.concatenate([objtype, reftype]))
+
+ # site filter
+ if types.ObjType.SITE in set(objtype):
+ raise NotImplementedError(
+ 'Contact sensor with site matching semantics not implemented for JAX'
+ ' backend.'
+ )
+
+ # body semantics
+ if types.ObjType.BODY in contact_sensor_type:
+ raise NotImplementedError(
+ 'Contact sensor with body matching semantics not implemented for JAX'
+ ' backend.'
+ )
+
+ # subtree semantics
+ if types.ObjType.XBODY in contact_sensor_type:
+ raise NotImplementedError(
+ 'Contact sensor with subtree matching semantics not implemented for'
+ ' JAX backend.'
+ )
+
+ # net force
+ if (m.sensor_intprm[is_contact_sensor, 1] == 3).any():
+ raise NotImplementedError(
+ 'Contact sensor with netforce reduction not implemented for JAX'
+ ' backend.'
+ )
+
mesh_geomid = set()
for g1, g2, ip in collision_driver.geom_pairs(m):
t1, t2 = m.geom_type[[g1, g2]]
diff --git a/mjx/mujoco/mjx/_src/io_test.py b/mjx/mujoco/mjx/_src/io_test.py
index e3ecd539..3edef119 100644
--- a/mjx/mujoco/mjx/_src/io_test.py
+++ b/mjx/mujoco/mjx/_src/io_test.py
@@ -247,6 +247,35 @@ class ModelIOTest(parameterized.TestCase):
np.array([0, 0, 1, 0, 1, 0, 0]),
)
+ @parameterized.parameters(
+ '',
+ '',
+ '',
+ '',
+ ''
+ '',
+ '',
+ '',
+ )
+ def test_contact_sensor_jax(self, contact_sensor):
+ m = mujoco.MjModel.from_xml_string(f"""
+
+
+
+
+
+
+
+
+
+
+ {contact_sensor}
+
+
+ """)
+ with self.assertRaises(NotImplementedError):
+ mjx.put_model(m, impl='jax')
+
class DataIOTest(parameterized.TestCase):
"""IO tests for mjx.Data."""
diff --git a/mjx/mujoco/mjx/_src/sensor.py b/mjx/mujoco/mjx/_src/sensor.py
index bbdab811..6ea327f3 100644
--- a/mjx/mujoco/mjx/_src/sensor.py
+++ b/mjx/mujoco/mjx/_src/sensor.py
@@ -456,6 +456,25 @@ def sensor_acc(m: Model, d: Data) -> Data:
}:
d = smooth.rne_postconstraint(m, d)
+ contact_intprm = m.sensor_intprm[m.sensor_type == SensorType.CONTACT]
+ contact_maxforce = (contact_intprm[:, 1] == 2).any()
+ contact_dataforce = (contact_intprm[:, 0] & (1 << 1)).any()
+ contact_datatorque = (contact_intprm[:, 0] & (1 << 2)).any()
+ if (m.sensor_type[stage_acc] == SensorType.TOUCH).any() | (
+ (m.sensor_type[stage_acc] == SensorType.CONTACT).any()
+ and (contact_maxforce | contact_dataforce | contact_datatorque)
+ ):
+ # compute contact forces
+ contact_force = []
+ condim_ids = []
+ for dim in set(d._impl.contact.dim):
+ force, condim_id = support.contact_force_dim(m, d, dim)
+ contact_force.append(force)
+ condim_ids.append(condim_id)
+ contact_force = jp.concatenate(contact_force)[
+ np.argsort(np.concatenate(condim_ids))
+ ]
+
sensors, adrs = [], []
for sensor_type in sensor_types:
@@ -466,15 +485,6 @@ def sensor_acc(m: Model, d: Data) -> Data:
data_type = m.sensor_datatype[idx]
if sensor_type == SensorType.TOUCH:
- # compute contact forces
- forces = []
- condim_ids = []
- for dim in set(d._impl.contact.dim):
- force, condim_id = support.contact_force_dim(m, d, dim)
- forces.append(force)
- condim_ids.append(condim_id)
- forces = jp.concatenate(forces)[np.argsort(np.concatenate(condim_ids))]
-
# get bodies of contact geoms
conbody = jp.array(m.geom_bodyid)[d._impl.contact.geom]
@@ -493,7 +503,7 @@ def sensor_acc(m: Model, d: Data) -> Data:
# compute conray, flip if second body
conray = jax.vmap(
lambda frame, force: math.normalize(frame[0] * force[0])
- )(d._impl.contact.frame, forces)
+ )(d._impl.contact.frame, contact_force)
conray = jp.where(conbody1[..., None], -conray, conray)
# compute distance, mapping over sites and contacts
@@ -523,7 +533,163 @@ def sensor_acc(m: Model, d: Data) -> Data:
dist = jp.vstack(dist)[np.argsort(np.concatenate(dist_id))]
# accumulate normal forces for each site
- sensor = jp.dot((dist > 0) & contacts, forces[:, 0])
+ sensor = jp.dot((dist > 0) & contacts, contact_force[:, 0])
+ elif sensor_type == SensorType.CONTACT:
+ # maximum number of contacts
+ ncon = d._impl.ncon
+
+ # active contacts
+ dist = d._impl.contact.dist
+ pos = dist - d._impl.contact.includemargin
+ is_contact = pos < 0
+
+ # reduction criteria
+ if contact_maxforce:
+ # compute force magnitude for each contact
+ force_mag = jax.vmap(
+ lambda forcetorque: jp.dot(forcetorque[:3], forcetorque[:3])
+ )(contact_force)
+
+ def _reduce(reduction, mask):
+ if reduction == 1: # mindist
+ return jp.argsort(pos * mask, descending=False)
+ if reduction == 2: # maxforce
+ return jp.argsort(force_mag * mask, descending=True)
+ return jp.arange(mask.size)
+
+ # number of data elements per slot
+ def nslotdata(dataspec):
+ size = 0
+ # found, force, torque, dist, pos, normal, tangent
+ # TODO(taylorhowell): get sizes from mjCONDATA_SIZE
+ for i, size_i in enumerate([1, 3, 3, 1, 3, 3, 3]):
+ if dataspec & (1 << i):
+ size += size_i
+ return size
+
+ dataspecs, reduces = m.sensor_intprm[idx].T
+ dims = m.sensor_dim[idx]
+ objtypes = m.sensor_objtype[idx]
+ refid = m.sensor_refid[idx]
+ reftypes = m.sensor_reftype[idx]
+
+ for dataspec, reduce, objtype, reftype, dim in set(
+ zip(dataspecs, reduces, objtypes, reftypes, dims)
+ ):
+ idx_ds = (
+ (dataspec == dataspecs)
+ & (reduce == reduces)
+ & (objtype == objtypes)
+ & (reftype == reftypes)
+ & (dim == dims)
+ )
+
+ # TODO(taylorhowell): site filter
+
+ size = nslotdata(dataspec)
+ num = np.minimum(int(dim / size), ncon)
+
+ if objtype == ObjType.UNKNOWN and reftype == ObjType.UNKNOWN:
+ # all contacts match
+ match = np.ones(ncon, dtype=np.bool)
+
+ # matched and reduced contact ids
+ sort = _reduce(reduce, match)
+ cid = sort[:num]
+
+ # number of contacts per sensor
+ nfound = sum(is_contact)
+
+ # if duplicate sensor
+ nsensor = idx_ds.sum()
+ cid = np.tile(cid, (nsensor,))
+ match = np.tile(match[:num], (nsensor,))
+ nfound = np.tile(nfound, (nsensor,))
+ flip = np.ones((cid.size, 3))
+ elif objtype == ObjType.GEOM or reftype == ObjType.GEOM:
+ sensorid1 = objid[idx_ds]
+ sensorid2 = refid[idx_ds]
+ geomid0 = d._impl.contact.geom[:, 0]
+ geomid1 = d._impl.contact.geom[:, 1]
+
+ # match sensor ids and contact geom ids
+ geom0id1 = geomid0 == sensorid1[:, None]
+ geom0id2 = geomid0 == sensorid2[:, None]
+ geom1id1 = geomid1 == sensorid1[:, None]
+ geom1id2 = geomid1 == sensorid2[:, None]
+
+ if objtype == ObjType.GEOM and reftype == ObjType.UNKNOWN: # geom1
+ mask12 = geom0id1
+ mask21 = geom1id1
+ elif objtype == ObjType.UNKNOWN and reftype == ObjType.GEOM: # geom2
+ mask12 = geom0id2
+ mask21 = geom1id2
+ else: # geom1, geom2
+ mask12 = geom0id1 & geom1id2
+ mask21 = geom0id2 & geom1id1
+
+ match = mask12 | mask21
+
+ # matched and reduced contact ids
+ cid = jax.vmap(lambda x: _reduce(reduce, x))(match)[:, :num]
+ cid = cid.reshape(-1)
+
+ # flip direction for force, torque, normal, tangent
+ if reftype == ObjType.UNKNOWN: # geom1
+ is_flip = (geomid1[cid] == np.repeat(sensorid1, num))[:, None]
+ elif objtype == ObjType.UNKNOWN: # geom2
+ is_flip = (geomid0[cid] == np.repeat(sensorid2, num))[:, None]
+ else: # geom1, geom2
+ is_flip = np.repeat(sensorid1 > sensorid2, num)[:, None]
+
+ flip = jp.where(
+ is_flip,
+ jp.array([[1, 1, -1]]),
+ jp.array([[1, 1, 1]]),
+ )
+
+ # number of contacts per sensor
+ nfound = (match * is_contact[None, :]).sum(axis=1)
+
+ match = match[:, :num].reshape(-1)
+
+ # TODO(taylorhowell): matching criteria: body, subtree
+
+ else:
+ raise NotImplementedError(
+ f'Unsupported contact sensor semantics: {objtype} {reftype}.'
+ )
+
+ slot = []
+
+ if dataspec & (1 << 0): # found
+ slot.append(jp.repeat(nfound, num)[:, None])
+
+ if dataspec & (1 << 1): # force
+ slot.append(flip * contact_force[cid, :3])
+
+ if dataspec & (1 << 2): # torque
+ slot.append(flip * contact_force[cid, 3:])
+
+ if dataspec & (1 << 3): # dist
+ slot.append(dist[cid, None])
+
+ if dataspec & (1 << 4): # pos
+ slot.append(d._impl.contact.pos[cid])
+
+ if dataspec & (1 << 5): # normal
+ slot.append(flip[:, 2, None] * d._impl.contact.frame[cid, 0])
+
+ if dataspec & (1 << 6): # tangent
+ slot.append(flip[:, 2, None] * d._impl.contact.frame[cid, 1])
+
+ found = is_contact[cid] & match
+ sensors.append((found[:, None] * np.hstack(slot)).reshape(-1))
+ adrs.append(
+ (adr[idx_ds][:, None] + np.arange(num * size)[None]).reshape(-1)
+ )
+ continue # avoid adding to sensors/adrs list a second time
+
elif sensor_type == SensorType.ACCELEROMETER:
@jax.vmap
diff --git a/mjx/mujoco/mjx/_src/sensor_test.py b/mjx/mujoco/mjx/_src/sensor_test.py
index cab50779..75677bab 100644
--- a/mjx/mujoco/mjx/_src/sensor_test.py
+++ b/mjx/mujoco/mjx/_src/sensor_test.py
@@ -16,6 +16,7 @@
from absl.testing import absltest
from absl.testing import parameterized
+import itertools
import jax
from jax import numpy as jp
import mujoco
@@ -97,6 +98,80 @@ class SensorTest(parameterized.TestCase):
# sensor values
_assert_eq(random_sensor, dx.sensordata, 'sensordata')
+ @parameterized.parameters(
+ 'type="sphere" size=".1"',
+ 'type="sphere" size=".05" margin=".045"',
+ 'type="capsule" size=".1 .1" euler="0 89 89"',
+ 'type="box" size=".1 .1 .1" euler=".05 .075 .1"',
+ )
+ def test_sensor_contact(self, geom):
+ """Tests contact sensor."""
+
+ field = ['found', 'force', 'torque', 'dist', 'pos', 'normal', 'tangent']
+ datas = itertools.chain.from_iterable(
+ itertools.combinations(field, i) for i in range(len(field) + 1)
+ )
+
+ contact_sensors = ''
+ for num in [1, 2, 3, 4, 5]:
+ for data in datas:
+ data = ' '.join(data)
+ for reduce in ['mindist', 'maxforce']:
+ for match in [
+ '',
+ '',
+ 'geom1="plane"',
+ 'geom1="geom1"',
+ 'geom1="sphere2"',
+ 'geom2="plane"',
+ 'geom2="geom1"',
+ 'geom2="sphere2"',
+ 'geom1="plane" geom2="geom1"',
+ 'geom1="geom1" geom2="plane"',
+ 'geom1="plane" geom2="sphere2"',
+ 'geom1="sphere2" geom2="plane"',
+ 'geom1="geom1" geom2="sphere2"',
+ 'geom1="sphere2" geom2="geom1"',
+ ]:
+ contact_sensors += (
+ f'\n'
+ )
+
+ _MJCF = f"""
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+ {contact_sensors}
+
+
+
+
+
+ """
+ m = mujoco.MjModel.from_xml_string(_MJCF)
+ d = mujoco.MjData(m)
+ mujoco.mj_resetDataKeyframe(m, d, 0)
+
+ mx = mjx.put_model(m)
+ dx = mjx.put_data(m, d)
+
+ mujoco.mj_forward(m, d)
+ dx = mjx.forward(mx, dx)
+
+ _assert_eq(dx.sensordata, d.sensordata, 'sensordata')
+
def test_unsupported_sensor(self):
"""Tests unsupported sensor raises error."""
m = mujoco.MjModel.from_xml_string("""
diff --git a/mjx/mujoco/mjx/_src/types.py b/mjx/mujoco/mjx/_src/types.py
index 42af00a7..19c84faa 100644
--- a/mjx/mujoco/mjx/_src/types.py
+++ b/mjx/mujoco/mjx/_src/types.py
@@ -380,6 +380,7 @@ class SensorType(enum.IntEnum):
SUBTREELINVEL: subtree linear velocity
SUBTREEANGMOM: subtree angular momentum
TOUCH: scalar contact normal forces summed over the sensor zone
+ CONTACT: contacts which occurred during the simulation
ACCELEROMETER: accelerometer
FORCE: force
TORQUE: torque
@@ -415,6 +416,7 @@ class SensorType(enum.IntEnum):
SUBTREELINVEL = mujoco.mjtSensor.mjSENS_SUBTREELINVEL
SUBTREEANGMOM = mujoco.mjtSensor.mjSENS_SUBTREEANGMOM
TOUCH = mujoco.mjtSensor.mjSENS_TOUCH
+ CONTACT = mujoco.mjtSensor.mjSENS_CONTACT
ACCELEROMETER = mujoco.mjtSensor.mjSENS_ACCELEROMETER
FORCE = mujoco.mjtSensor.mjSENS_FORCE
TORQUE = mujoco.mjtSensor.mjSENS_TORQUE
@@ -577,7 +579,6 @@ class ModelC(PyTreeNode):
flex_bvhnum: jax.Array
actuator_plugin: jax.Array
sensor_plugin: jax.Array
- sensor_intprm: jax.Array
plugin: jax.Array
plugin_stateadr: jax.Array
@@ -843,6 +844,7 @@ class Model(PyTreeNode):
sensor_objid: np.ndarray
sensor_reftype: np.ndarray
sensor_refid: np.ndarray
+ sensor_intprm: np.ndarray
sensor_dim: np.ndarray
sensor_adr: np.ndarray
sensor_cutoff: np.ndarray