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