From b2174a7ec30dbd0c8dc554699e9d3c4dc4731233 Mon Sep 17 00:00:00 2001 From: Taylor Howell Date: Thu, 22 Aug 2024 02:20:20 -0700 Subject: [PATCH] Add framequat sensor to MJX. PiperOrigin-RevId: 666254481 Change-Id: I50be8574e1e4677ef5dc23a2bdcbc296f301d7fc --- doc/changelog.rst | 4 +- mjx/mujoco/mjx/_src/sensor.py | 94 ++++++++++++++++++++--------- mjx/mujoco/mjx/_src/types.py | 2 + mjx/mujoco/mjx/test_data/sensor.xml | 4 ++ 4 files changed, 72 insertions(+), 32 deletions(-) diff --git a/doc/changelog.rst b/doc/changelog.rst index 695122ee..195d20cd 100644 --- a/doc/changelog.rst +++ b/doc/changelog.rst @@ -18,8 +18,8 @@ MJX ^^^ 5. Added ``efc_pos`` to ``mjx.Data`` (:github:issue:`1388`). 6. Added position-dependent sensors: ``MAGNETOMETER``, ``CAMPROJECTION``, ``RANGEFINDER``, ``JOINTPOS``, - ``ACTUATORPOS``, ``BALLQUAT``, ``FRAMEPOS``, ``FRAMEXAXIS``, ``FRAMEYAXIS``, ``FRAMEZAXIS``, ``SUBTREECOM``, - ``CLOCK``. + ``ACTUATORPOS``, ``BALLQUAT``, ``FRAMEPOS``, ``FRAMEXAXIS``, ``FRAMEYAXIS``, ``FRAMEZAXIS``, ``FRAMEQUAT``, + ``SUBTREECOM``, ``CLOCK``. 7. Added velocity-dependent sensors: ``JOINTVEL``, ``ACTUATORVEL``, ``BALLANGVEL``. 8. Added acceleration/force-dependent sensors: ``ACTUATORFRC``, ``JOINTACTFRC``. 9. Changed default policy to avoid placing unused (MuJoCo-only) arrays on device. diff --git a/mjx/mujoco/mjx/_src/sensor.py b/mjx/mujoco/mjx/_src/sensor.py index 4e704472..d2ad1497 100644 --- a/mjx/mujoco/mjx/_src/sensor.py +++ b/mjx/mujoco/mjx/_src/sensor.py @@ -61,6 +61,9 @@ def sensor_pos(m: Model, d: Data) -> Data: for sensor_type in set(m.sensor_type[stage_pos]): idx = m.sensor_type == sensor_type objid = m.sensor_objid[idx] + objtype = m.sensor_objtype[idx] + refid = m.sensor_refid[idx] + reftype = m.sensor_reftype[idx] adr = m.sensor_adr[idx] if sensor_type == SensorType.MAGNETOMETER: @@ -113,7 +116,6 @@ def sensor_pos(m: Model, d: Data) -> Data: return sensor[:2] - refid = m.sensor_refid[idx] sensorsize = m.cam_sensorsize[refid] intrinsic = m.cam_intrinsic[refid] fovy = m.cam_fovy[refid] @@ -131,15 +133,15 @@ def sensor_pos(m: Model, d: Data) -> Data: 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])] + idxs = sid == site_bodyid + objids = objid[idxs] + site_xpos = d.site_xpos[objids] + site_mat = d.site_xmat[objids].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_]) + adrs.append(adr[idxs]) continue # avoid adding to sensors/adrs list a second time elif sensor_type == SensorType.JOINTPOS: sensor = d.qpos[m.jnt_qposadr[objid]] @@ -155,23 +157,19 @@ def sensor_pos(m: Model, d: Data) -> Data: def _framepos(xpos, xpos_ref, xmat_ref, refid): return jp.where(refid == -1, xpos, xmat_ref.T @ (xpos - xpos_ref)) - objtype = m.sensor_objtype[idx] - reftype = m.sensor_reftype[idx] - refid = m.sensor_refid[idx] - # evaluate for valid object and reference object type pairs for ot, rt in set(zip(objtype, reftype)): - id_ = (objtype == ot) & (reftype == rt) - refid_ = refid[id_] + idxt = (objtype == ot) & (reftype == rt) + refidt = refid[idxt] xpos, _ = objtype_data[ot] xpos_ref, xmat_ref = objtype_data[rt] - xpos = xpos[objid[id_]] - xpos_ref = xpos_ref[refid_] - xmat_ref = xmat_ref[refid_] - sensor = jax.vmap(_framepos)(xpos, xpos_ref, xmat_ref, refid_) - adr_ = adr[id_, None] + np.arange(3)[None] + xpos = xpos[objid[idxt]] + xpos_ref = xpos_ref[refidt] + xmat_ref = xmat_ref[refidt] + sensor = jax.vmap(_framepos)(xpos, xpos_ref, xmat_ref, refidt) + adrt = adr[idxt, None] + np.arange(3)[None] sensors.append(sensor.reshape(-1)) - adrs.append(adr_.reshape(-1)) + adrs.append(adrt.reshape(-1)) continue # avoid adding to sensors/adrs list a second time elif sensor_type in frame_axis: @@ -179,22 +177,58 @@ def sensor_pos(m: Model, d: Data) -> Data: axis = xmat[:, frame_axis[sensor_type]] return jp.where(refid == -1, axis, xmat_ref.T @ axis) - objtype = m.sensor_objtype[idx] - reftype = m.sensor_reftype[idx] - refid = m.sensor_refid[idx] + # evaluate for valid object and reference object type pairs + for ot, rt in set(zip(objtype, reftype)): + idxt = (objtype == ot) & (reftype == rt) + refidt = refid[idxt] + _, xmat = objtype_data[ot] + _, xmat_ref = objtype_data[rt] + xmat = xmat[objid[idxt]] + xmat_ref = xmat_ref[refidt] + sensor = jax.vmap(_frameaxis)(xmat, xmat_ref, refidt) + adrt = adr[idxt, None] + np.arange(3)[None] + sensors.append(sensor.reshape(-1)) + adrs.append(adrt.reshape(-1)) + continue # avoid adding to sensors/adrs list a second time + elif sensor_type == SensorType.FRAMEQUAT: + + def _quat(otype, oid): + if otype == ObjType.XBODY: + return d.xquat[oid] + elif otype == ObjType.BODY: + return jax.vmap(math.quat_mul)(d.xquat[oid], m.body_iquat[oid]) + elif otype == ObjType.GEOM: + return jax.vmap(math.quat_mul)( + d.xquat[m.geom_bodyid[oid]], m.geom_quat[oid] + ) + elif otype == ObjType.SITE: + return jax.vmap(math.quat_mul)( + d.xquat[m.site_bodyid[oid]], m.site_quat[oid] + ) + elif otype == ObjType.CAMERA: + return jax.vmap(math.quat_mul)( + d.xquat[m.cam_bodyid[oid]], m.cam_quat[oid] + ) + elif otype == ObjType.UNKNOWN: + return jp.tile(jp.array([1.0, 0.0, 0.0, 0.0]), (oid.size, 1)) + else: + raise ValueError(f'Unknown object type: {otype}') # evaluate for valid object and reference object type pairs for ot, rt in set(zip(objtype, reftype)): - id_ = (objtype == ot) & (reftype == rt) - refid_ = refid[id_] - _, xmat = objtype_data[ot] - _, xmat_ref = objtype_data[rt] - xmat = xmat[objid[id_]] - xmat_ref = xmat_ref[refid_] - sensor = jax.vmap(_frameaxis)(xmat, xmat_ref, refid_) - adr_ = adr[id_, None] + np.arange(3)[None] + idxt = (objtype == ot) & (reftype == rt) + objidt = objid[idxt] + refidt = refid[idxt] + quat = _quat(ot, objidt) + refquat = _quat(rt, refidt) + sensor = jax.vmap( + lambda q, r, rid: jp.where( + rid == -1, q, math.quat_mul(math.quat_inv(r), q) + ) + )(quat, refquat, refidt) + adrt = adr[idxt, None] + np.arange(4)[None] sensors.append(sensor.reshape(-1)) - adrs.append(adr_.reshape(-1)) + adrs.append(adrt.reshape(-1)) continue # avoid adding to sensors/adrs list a second time elif sensor_type == SensorType.SUBTREECOM: sensor = d.subtree_com[objid].reshape(-1) diff --git a/mjx/mujoco/mjx/_src/types.py b/mjx/mujoco/mjx/_src/types.py index fd3aa362..eaf54fd4 100644 --- a/mjx/mujoco/mjx/_src/types.py +++ b/mjx/mujoco/mjx/_src/types.py @@ -304,6 +304,7 @@ class SensorType(enum.IntEnum): FRAMEXAXIS: frame x-axis FRAMEYAXIS: frame y-axis FRAMEZAXIS: frame z-axis + FRAMEQUAT: frame orientation, represented as quaternion SUBTREECOM: subtree centor of mass CLOCK: simulation time JOINTVEL: joint velocity @@ -322,6 +323,7 @@ class SensorType(enum.IntEnum): FRAMEXAXIS = mujoco.mjtSensor.mjSENS_FRAMEXAXIS FRAMEYAXIS = mujoco.mjtSensor.mjSENS_FRAMEYAXIS FRAMEZAXIS = mujoco.mjtSensor.mjSENS_FRAMEZAXIS + FRAMEQUAT = mujoco.mjtSensor.mjSENS_FRAMEQUAT SUBTREECOM = mujoco.mjtSensor.mjSENS_SUBTREECOM CLOCK = mujoco.mjtSensor.mjSENS_CLOCK JOINTVEL = mujoco.mjtSensor.mjSENS_JOINTVEL diff --git a/mjx/mujoco/mjx/test_data/sensor.xml b/mjx/mujoco/mjx/test_data/sensor.xml index 80285609..5a636e31 100644 --- a/mjx/mujoco/mjx/test_data/sensor.xml +++ b/mjx/mujoco/mjx/test_data/sensor.xml @@ -11,6 +11,7 @@ -framexaxis -frameyaxis -framezaxis +-framequat -subtreecom -clock * velocity-dependent sensors: @@ -50,6 +51,7 @@ + @@ -91,6 +93,7 @@ + @@ -109,6 +112,7 @@ +