From 6db96e07e43f441a89cdc129ac751b6d8d7567ef Mon Sep 17 00:00:00 2001 From: Taylor Howell Date: Fri, 16 Aug 2024 03:09:23 -0700 Subject: [PATCH] Add position-dependent sensors to MJX: magnetometer, ballquat, subtreecom, framepos, framexaxis, frameyaxis, framezaxis, and clock. PiperOrigin-RevId: 663668212 Change-Id: Ifc9adc6bff4544172ab22572da04a093aebfdbd4 --- doc/changelog.rst | 2 + mjx/mujoco/mjx/_src/sensor.py | 121 ++++++++++++++++++++++++---- mjx/mujoco/mjx/_src/types.py | 35 ++++++++ mjx/mujoco/mjx/test_data/sensor.xml | 56 +++++++++++-- 4 files changed, 195 insertions(+), 19 deletions(-) diff --git a/doc/changelog.rst b/doc/changelog.rst index 8a29045c..9504fd2d 100644 --- a/doc/changelog.rst +++ b/doc/changelog.rst @@ -14,6 +14,8 @@ General MJX ^^^ 4. Added ``efc_pos`` to ``mjx.Data``. +5. Added position-dependent sensors: ``MAGNETOMETER``, ``JOINTPOS``, ``ACTUATORPOS``, ``BALLQUAT``, ``FRAMEPOS``, + ``FRAMEXAXIS``, ``FRAMEYAXIS``, ``FRAMEZAXIS``, ``SUBTREECOM``, ``CLOCK``. Version 3.2.2 (Aug 8, 2024) --------------------------- diff --git a/mjx/mujoco/mjx/_src/sensor.py b/mjx/mujoco/mjx/_src/sensor.py index a59e9c58..b92c32e0 100644 --- a/mjx/mujoco/mjx/_src/sensor.py +++ b/mjx/mujoco/mjx/_src/sensor.py @@ -15,11 +15,14 @@ """Sensor functions.""" import jax +from jax import numpy as jp +import mujoco # pylint: disable=g-importing-member +from mujoco.mjx._src import math from mujoco.mjx._src.types import Data from mujoco.mjx._src.types import Model +from mujoco.mjx._src.types import ObjType from mujoco.mjx._src.types import SensorType -from typing import Tuple # pylint: enable=g-importing-member import numpy as np @@ -27,19 +30,109 @@ import numpy as np def sensor_pos(m: Model, d: Data) -> Data: """Compute position-dependent sensors values.""" - sensordata = d.sensordata - if np.isin(SensorType.JOINTPOS, m.sensor_type): - # jointpos - i = m.sensor_type == SensorType.JOINTPOS - objid = m.sensor_objid[i] - adr = m.sensor_adr[i] - sensordata = sensordata.at[adr].set(d.qpos[m.jnt_qposadr[objid]]) - if np.isin(SensorType.ACTUATORPOS, m.sensor_type): - # actuatorpos - i = m.sensor_type == SensorType.ACTUATORPOS - objid = m.sensor_objid[i] - adr = m.sensor_adr[i] - sensordata = sensordata.at[adr].set(d.actuator_length[objid]) + # no position-dependent sensors + if sum(m.sensor_needstage == mujoco.mjtStage.mjSTAGE_POS) == 0: + return d + + # position and orientation by object type + objtype_data = { + ObjType.UNKNOWN: ( + np.expand_dims(np.eye(3), axis=0), + np.zeros((1, 3)), + ), # world + ObjType.BODY: (d.xipos, d.ximat), + ObjType.XBODY: (d.xpos, d.xmat), + ObjType.GEOM: (d.geom_xpos, d.geom_xmat), + ObjType.SITE: (d.site_xpos, d.site_xmat), + ObjType.CAMERA: (d.cam_xpos, d.cam_xmat), + } + + # frame axis indexing + frame_axis = { + SensorType.FRAMEXAXIS: 0, + SensorType.FRAMEYAXIS: 1, + SensorType.FRAMEZAXIS: 2, + } + + sensors, adrs = [], [] + + for sensor_type in set(m.sensor_type): + idx = m.sensor_type == sensor_type + objid = m.sensor_objid[idx] + adr = m.sensor_adr[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.JOINTPOS: + sensor = d.qpos[m.jnt_qposadr[objid]] + elif sensor_type == SensorType.ACTUATORPOS: + sensor = d.actuator_length[objid] + 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) + adr = (adr[:, None] + np.arange(4)[None]).reshape(-1) + elif sensor_type == SensorType.FRAMEPOS: + + 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_] + 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] + sensors.append(sensor.reshape(-1)) + adrs.append(adr_.reshape(-1)) + continue # avoid adding to sensors/adrs list a second time + elif sensor_type in frame_axis: + + def _frameaxis(xmat, xmat_ref, refid): + 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)): + 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] + sensors.append(sensor.reshape(-1)) + adrs.append(adr_.reshape(-1)) + continue # avoid adding to sensors/adrs list a second time + elif sensor_type == SensorType.SUBTREECOM: + sensor = d.subtree_com[objid].reshape(-1) + adr = (adr[:, None] + np.arange(3)[None]).reshape(-1) + elif sensor_type == SensorType.CLOCK: + sensor = jp.repeat(d.time, sum(idx)) + + sensors.append(sensor) + adrs.append(adr) + + sensordata = d.sensordata.at[np.concatenate(adrs)].set( + jp.concatenate(sensors) + ) return d.replace(sensordata=sensordata) diff --git a/mjx/mujoco/mjx/_src/types.py b/mjx/mujoco/mjx/_src/types.py index 568cec28..8b690c11 100644 --- a/mjx/mujoco/mjx/_src/types.py +++ b/mjx/mujoco/mjx/_src/types.py @@ -282,11 +282,46 @@ class SensorType(enum.IntEnum): """Type of sensor. Members: + MAGNETOMETER: magnetometer JOINTPOS: joint position ACTUATORPOS: actuator position + BALLQUAT: ball joint orientation + FRAMEPOS: frame position + FRAMEXAXIS: frame x-axis + FRAMEYAXIS: frame y-axis + FRAMEZAXIS: frame z-axis + SUBTREECOM: subtree centor of mass + CLOCK: simulation time """ + MAGNETOMETER = mujoco.mjtSensor.mjSENS_MAGNETOMETER JOINTPOS = mujoco.mjtSensor.mjSENS_JOINTPOS ACTUATORPOS = mujoco.mjtSensor.mjSENS_ACTUATORPOS + BALLQUAT = mujoco.mjtSensor.mjSENS_BALLQUAT + FRAMEPOS = mujoco.mjtSensor.mjSENS_FRAMEPOS + FRAMEXAXIS = mujoco.mjtSensor.mjSENS_FRAMEXAXIS + FRAMEYAXIS = mujoco.mjtSensor.mjSENS_FRAMEYAXIS + FRAMEZAXIS = mujoco.mjtSensor.mjSENS_FRAMEZAXIS + SUBTREECOM = mujoco.mjtSensor.mjSENS_SUBTREECOM + CLOCK = mujoco.mjtSensor.mjSENS_CLOCK + + +class ObjType(PyTreeNode): + """Type of object. + + Members: + UNKNOWN: unknown object type + BODY: body + XBODY: body, used to access regular frame instead of i-frame + GEOM: geom + SITE: site + CAMERA: camera + """ + UNKNOWN = mujoco.mjtObj.mjOBJ_UNKNOWN + BODY = mujoco.mjtObj.mjOBJ_BODY + XBODY = mujoco.mjtObj.mjOBJ_XBODY + GEOM = mujoco.mjtObj.mjOBJ_GEOM + SITE = mujoco.mjtObj.mjOBJ_SITE + CAMERA = mujoco.mjtObj.mjOBJ_CAMERA class Option(PyTreeNode): diff --git a/mjx/mujoco/mjx/test_data/sensor.xml b/mjx/mujoco/mjx/test_data/sensor.xml index 810abd1d..13178604 100644 --- a/mjx/mujoco/mjx/test_data/sensor.xml +++ b/mjx/mujoco/mjx/test_data/sensor.xml @@ -1,28 +1,74 @@ - + + + + + + + + + + + + + + + + + + + - + + + + + + + + + + + + + + + + + + +