Add sensor cutoff to MJX.

PiperOrigin-RevId: 669363779
Change-Id: Ia5c3f913c1b84e1893aa6072268a334123f00aeb
This commit is contained in:
Taylor Howell
2024-08-30 10:08:19 -07:00
committed by Copybara-Service
parent cf02171bf8
commit 5391e6011a
2 changed files with 45 additions and 14 deletions
+41 -14
View File
@@ -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:
+4
View File
@@ -85,10 +85,14 @@
<sensor>
<magnetometer name="magnetometer0" site="site0"/>
<velocimeter name="velocimeter0" site="site0"/>
<velocimeter name="velocimeter0cutoff" site="site0" cutoff="3e-4"/>
<gyro name="gyro0" site="site0"/>
<gyro name="gyro0cutoff" site="site0" cutoff="2e-3"/>
<rangefinder name="rangefinder0" site="site_rangefinder0"/>
<jointpos name="jointpos0" joint="hinge0"/>
<jointpos name="jointpos0cutoff" joint="hinge0" cutoff="1e-4"/>
<jointvel name="jointvel0" joint="hinge0"/>
<jointvel name="jointvel0cutoff" joint="hinge0" cutoff="1e-3"/>
<actuatorfrc name="actuatorfrc0" actuator="motor0"/>
<actuatorpos name="actuatorpos0" actuator="motor0"/>
<actuatorvel name="actuatorvel0" actuator="motor0"/>