From 5391e6011a3f77af146af13807dbb721ed92bb19 Mon Sep 17 00:00:00 2001 From: Taylor Howell Date: Fri, 30 Aug 2024 10:08:19 -0700 Subject: [PATCH] Add sensor cutoff to MJX. PiperOrigin-RevId: 669363779 Change-Id: Ia5c3f913c1b84e1893aa6072268a334123f00aeb --- mjx/mujoco/mjx/_src/sensor.py | 55 +++++++++++++++++++++-------- mjx/mujoco/mjx/test_data/sensor.xml | 4 +++ 2 files changed, 45 insertions(+), 14 deletions(-) diff --git a/mjx/mujoco/mjx/_src/sensor.py b/mjx/mujoco/mjx/_src/sensor.py index bb090496..e5cbf604 100644 --- a/mjx/mujoco/mjx/_src/sensor.py +++ b/mjx/mujoco/mjx/_src/sensor.py @@ -29,6 +29,23 @@ from mujoco.mjx._src.types import SensorType import numpy as np +def apply_cutoff( + sensor: jax.Array, cutoff: jax.Array, data_type: int +) -> jax.Array: + """Clip sensor to cutoff value.""" + + @jax.vmap + def fn(sensor, cutoff): + if data_type == mujoco.mjtDataType.mjDATATYPE_REAL: + return jp.where(cutoff > 0, jp.clip(sensor, -cutoff, cutoff), sensor) + elif data_type == mujoco.mjtDataType.mjDATATYPE_POSITIVE: + return jp.where(cutoff > 0, jp.minimum(sensor, cutoff), sensor) + else: + return sensor + + return fn(sensor, cutoff) + + def sensor_pos(m: Model, d: Data) -> Data: """Compute position-dependent sensors values.""" @@ -65,11 +82,13 @@ def sensor_pos(m: Model, d: Data) -> Data: refid = m.sensor_refid[idx] reftype = m.sensor_reftype[idx] adr = m.sensor_adr[idx] + cutoff = m.sensor_cutoff[idx] + data_type = m.sensor_datatype[idx] if sensor_type == SensorType.MAGNETOMETER: sensor = jax.vmap(lambda xmat: xmat.T @ m.opt.magnetic)( d.site_xmat[objid] - ).reshape(-1) + ) adr = (adr[:, None] + np.arange(3)[None]).reshape(-1) elif sensor_type == SensorType.CAMPROJECTION: @@ -128,7 +147,7 @@ def sensor_pos(m: Model, d: Data) -> Data: sensor = _cam_project( target_xpos, xpos, xmat, res, fovy, intrinsic, sensorsize, focal_flag - ).reshape(-1) + ) adr = (adr[:, None] + np.arange(2)[None]).reshape(-1) elif sensor_type == SensorType.RANGEFINDER: site_bodyid = m.site_bodyid[objid] @@ -137,10 +156,11 @@ def sensor_pos(m: Model, d: Data) -> Data: objids = objid[idxs] site_xpos = d.site_xpos[objids] site_mat = d.site_xmat[objids].reshape((-1, 9))[:, np.array([2, 5, 8])] + cutoffs = cutoff[idxs] 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) + sensors.append(apply_cutoff(sensor, cutoffs, data_type[0])) adrs.append(adr[idxs]) continue # avoid adding to sensors/adrs list a second time elif sensor_type == SensorType.JOINTPOS: @@ -150,7 +170,7 @@ def sensor_pos(m: Model, d: Data) -> Data: elif sensor_type == SensorType.BALLQUAT: jnt_qposadr = m.jnt_qposadr[objid, None] + np.arange(4)[None] quat = d.qpos[jnt_qposadr] - sensor = jax.vmap(math.normalize)(quat).reshape(-1) + sensor = jax.vmap(math.normalize)(quat) adr = (adr[:, None] + np.arange(4)[None]).reshape(-1) elif sensor_type == SensorType.FRAMEPOS: @@ -166,9 +186,10 @@ def sensor_pos(m: Model, d: Data) -> Data: xpos = xpos[objid[idxt]] xpos_ref = xpos_ref[refidt] xmat_ref = xmat_ref[refidt] + cutofft = cutoff[idxt] sensor = jax.vmap(_framepos)(xpos, xpos_ref, xmat_ref, refidt) adrt = adr[idxt, None] + np.arange(3)[None] - sensors.append(sensor.reshape(-1)) + sensors.append(apply_cutoff(sensor, cutofft, data_type[0]).reshape(-1)) adrs.append(adrt.reshape(-1)) continue # avoid adding to sensors/adrs list a second time elif sensor_type in frame_axis: @@ -185,9 +206,10 @@ def sensor_pos(m: Model, d: Data) -> Data: _, xmat_ref = objtype_data[rt] xmat = xmat[objid[idxt]] xmat_ref = xmat_ref[refidt] + cutofft = cutoff[idxt] sensor = jax.vmap(_frameaxis)(xmat, xmat_ref, refidt) adrt = adr[idxt, None] + np.arange(3)[None] - sensors.append(sensor.reshape(-1)) + sensors.append(apply_cutoff(sensor, cutofft, data_type[0]).reshape(-1)) adrs.append(adrt.reshape(-1)) continue # avoid adding to sensors/adrs list a second time elif sensor_type == SensorType.FRAMEQUAT: @@ -221,17 +243,18 @@ def sensor_pos(m: Model, d: Data) -> Data: refidt = refid[idxt] quat = _quat(ot, objidt) refquat = _quat(rt, refidt) + cutofft = cutoff[idxt] 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)) + sensors.append(apply_cutoff(sensor, cutofft, data_type[0]).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) + sensor = d.subtree_com[objid] adr = (adr[:, None] + np.arange(3)[None]).reshape(-1) elif sensor_type == SensorType.CLOCK: sensor = jp.repeat(d.time, sum(idx)) @@ -239,7 +262,7 @@ def sensor_pos(m: Model, d: Data) -> Data: # TODO(taylorhowell): raise error after adding sensor check to io.py continue # unsupported sensor type - sensors.append(sensor) + sensors.append(apply_cutoff(sensor, cutoff, data_type[0]).reshape(-1)) adrs.append(adr) if not adrs: @@ -265,6 +288,8 @@ def sensor_vel(m: Model, d: Data) -> Data: idx = m.sensor_type == sensor_type objid = m.sensor_objid[idx] adr = m.sensor_adr[idx] + cutoff = m.sensor_cutoff[idx] + data_type = m.sensor_datatype[idx] if sensor_type == SensorType.VELOCIMETER: bodyid = m.site_bodyid[objid] @@ -274,13 +299,13 @@ def sensor_vel(m: Model, d: Data) -> Data: subtree_com = d.subtree_com[m.body_rootid[bodyid]] sensor = jax.vmap( lambda vec, dif, rot: rot.T @ (vec[3:] - jp.cross(dif, vec[:3])) - )(cvel, pos - subtree_com, rot).reshape(-1) + )(cvel, pos - subtree_com, rot) adr = (adr[:, None] + np.arange(3)[None]).reshape(-1) elif sensor_type == SensorType.GYRO: bodyid = m.site_bodyid[objid] rot = d.site_xmat[objid] ang = d.cvel[bodyid, :3] - sensor = jax.vmap(lambda ang, rot: rot.T @ ang)(ang, rot).reshape(-1) + sensor = jax.vmap(lambda ang, rot: rot.T @ ang)(ang, rot) adr = (adr[:, None] + np.arange(3)[None]).reshape(-1) elif sensor_type == SensorType.JOINTVEL: sensor = d.qvel[m.jnt_dofadr[objid]] @@ -288,13 +313,13 @@ def sensor_vel(m: Model, d: Data) -> Data: sensor = d.actuator_velocity[objid] elif sensor_type == SensorType.BALLANGVEL: jnt_dotadr = m.jnt_dofadr[objid, None] + np.arange(3)[None] - sensor = d.qvel[jnt_dotadr].reshape(-1) + sensor = d.qvel[jnt_dotadr] adr = (adr[:, None] + np.arange(3)[None]).reshape(-1) else: # TODO(taylorhowell): raise error after adding sensor check to io.py continue # unsupported sensor typ - sensors.append(sensor) + sensors.append(apply_cutoff(sensor, cutoff, data_type[0]).reshape(-1)) adrs.append(adr) if not adrs: @@ -320,6 +345,8 @@ def sensor_acc(m: Model, d: Data) -> Data: idx = m.sensor_type == sensor_type objid = m.sensor_objid[idx] adr = m.sensor_adr[idx] + cutoff = m.sensor_cutoff[idx] + data_type = m.sensor_datatype[idx] if sensor_type == SensorType.ACTUATORFRC: sensor = d.actuator_force[objid] @@ -329,7 +356,7 @@ def sensor_acc(m: Model, d: Data) -> Data: # TODO(taylorhowell): raise error after adding sensor check to io.py continue # unsupported sensor type - sensors.append(sensor) + sensors.append(apply_cutoff(sensor, cutoff, data_type[0]).reshape(-1)) adrs.append(adr) if not adrs: diff --git a/mjx/mujoco/mjx/test_data/sensor.xml b/mjx/mujoco/mjx/test_data/sensor.xml index c1231457..d1417b31 100644 --- a/mjx/mujoco/mjx/test_data/sensor.xml +++ b/mjx/mujoco/mjx/test_data/sensor.xml @@ -85,10 +85,14 @@ + + + +