Add framelinvel and frameangvel sensors to MJX.

PiperOrigin-RevId: 669647542
Change-Id: I612230e81b456c26f29173d1cb5a60975791875a
This commit is contained in:
Taylor Howell
2024-08-31 06:11:57 -07:00
committed by Copybara-Service
parent 651868ab05
commit 9805df616b
4 changed files with 91 additions and 4 deletions
+2 -1
View File
@@ -32,7 +32,8 @@ MJX
- Added position-dependent sensors: ``MAGNETOMETER``, ``CAMPROJECTION``, ``RANGEFINDER``, ``JOINTPOS``,
``ACTUATORPOS``, ``BALLQUAT``, ``FRAMEPOS``, ``FRAMEXAXIS``, ``FRAMEYAXIS``, ``FRAMEZAXIS``, ``FRAMEQUAT``,
``SUBTREECOM``, ``CLOCK``.
- Added velocity-dependent sensors: ``VELOCIMETER``, ``GYRO``, ``JOINTVEL``, ``ACTUATORVEL``, ``BALLANGVEL``.
- Added velocity-dependent sensors: ``VELOCIMETER``, ``GYRO``, ``JOINTVEL``, ``ACTUATORVEL``, ``BALLANGVEL``,
``FRAMELINVEL``, ``FRAMEANGVEL``.
- Added acceleration/force-dependent sensors: ``ACTUATORFRC``, ``JOINTACTFRC``.
- Changed default policy to avoid placing unused (MuJoCo-only) arrays on device.
- Added ``device`` parameter to ``mjx.make_data`` to bring it to parity with ``mjx.put_model`` and ``mjx.put_data``.
+72 -2
View File
@@ -55,8 +55,8 @@ def sensor_pos(m: Model, d: Data) -> Data:
# position and orientation by object type
objtype_data = {
ObjType.UNKNOWN: (
np.expand_dims(np.eye(3), axis=0),
np.zeros((1, 3)),
np.expand_dims(np.eye(3), axis=0),
), # world
ObjType.BODY: (d.xipos, d.ximat),
ObjType.XBODY: (d.xpos, d.xmat),
@@ -281,6 +281,20 @@ def sensor_vel(m: Model, d: Data) -> Data:
if m.opt.disableflags & DisableBit.SENSOR:
return d
# position and orientation by object type
objtype_data = {
ObjType.UNKNOWN: (
np.zeros((1, 3)),
np.expand_dims(np.eye(3), axis=0),
np.arange(1),
), # world
ObjType.BODY: (d.xipos, d.ximat, np.arange(m.nbody)),
ObjType.XBODY: (d.xpos, d.xmat, np.arange(m.nbody)),
ObjType.GEOM: (d.geom_xpos, d.geom_xmat, m.geom_bodyid),
ObjType.SITE: (d.site_xpos, d.site_xmat, m.site_bodyid),
ObjType.CAMERA: (d.cam_xpos, d.cam_xmat, m.cam_bodyid),
}
stage_vel = m.sensor_needstage == mujoco.mjtStage.mjSTAGE_VEL
sensors, adrs = [], []
@@ -315,9 +329,65 @@ def sensor_vel(m: Model, d: Data) -> Data:
jnt_dotadr = m.jnt_dofadr[objid, None] + np.arange(3)[None]
sensor = d.qvel[jnt_dotadr]
adr = (adr[:, None] + np.arange(3)[None]).reshape(-1)
elif sensor_type in (SensorType.FRAMELINVEL, SensorType.FRAMEANGVEL):
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)
objidt = objid[idxt]
refidt = refid[idxt]
cutofft = cutoff[idxt]
xpos, _, _ = objtype_data[ot]
xposref, xmatref, _ = objtype_data[rt]
xpos = xpos[objidt]
xposref = xposref[refidt]
xmatref = xmatref[refidt]
def _cvel_offset(otype, oid):
pos, _, bodyid = objtype_data[otype]
pos = pos[oid]
bodyid = bodyid[oid]
return d.cvel[bodyid], pos - d.subtree_com[m.body_rootid[bodyid]]
cvel, offset = _cvel_offset(ot, objidt)
cvelref, offsetref = _cvel_offset(rt, refidt)
cangvel = cvel[:, :3]
cangvelref = cvelref[:, :3]
if sensor_type == SensorType.FRAMELINVEL:
clinvel = cvel[:, 3:]
clinvelref = cvelref[:, 3:]
xlinvel = clinvel - jp.cross(offset, cangvel)
xlinvelref = clinvelref - jp.cross(offsetref, cangvelref)
rvec = xpos - xposref
rel_vel = xlinvel - xlinvelref + jp.cross(rvec, cangvelref)
sensor = jp.where(
(refidt > -1)[:, None],
jax.vmap(lambda mat, vec: mat.T @ vec)(xmatref, rel_vel),
xlinvel,
)
elif sensor_type == SensorType.FRAMEANGVEL:
rel_vel = cangvel - cangvelref
sensor = jp.where(
(refidt > -1)[:, None],
jax.vmap(lambda mat, vec: mat.T @ vec)(xmatref, rel_vel),
cangvel,
)
else:
raise ValueError(f'Unknown sensor type: {sensor_type}')
adrt = adr[idxt, None] + np.arange(3)[None]
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
else:
# TODO(taylorhowell): raise error after adding sensor check to io.py
continue # unsupported sensor typ
continue # unsupported sensor type
sensors.append(apply_cutoff(sensor, cutoff, data_type[0]).reshape(-1))
adrs.append(adr)
+4
View File
@@ -312,6 +312,8 @@ class SensorType(enum.IntEnum):
JOINTVEL: joint velocity
ACTUATORVEL: actuator velocity
BALLANGVEL: ball joint angular velocity
FRAMELINVEL: 3D linear velocity
FRAMEANGVEL: 3D angular velocity
ACTUATORFRC: scalar actuator force
JOINTACTFRC: scalar actuator force, measured at the joint
"""
@@ -333,6 +335,8 @@ class SensorType(enum.IntEnum):
JOINTVEL = mujoco.mjtSensor.mjSENS_JOINTVEL
ACTUATORVEL = mujoco.mjtSensor.mjSENS_ACTUATORVEL
BALLANGVEL = mujoco.mjtSensor.mjSENS_BALLANGVEL
FRAMELINVEL = mujoco.mjtSensor.mjSENS_FRAMELINVEL
FRAMEANGVEL = mujoco.mjtSensor.mjSENS_FRAMEANGVEL
ACTUATORFRC = mujoco.mjtSensor.mjSENS_ACTUATORFRC
JOINTACTFRC = mujoco.mjtSensor.mjSENS_JOINTACTFRC
+13 -1
View File
@@ -20,6 +20,8 @@
-jointvel
-actuatorvel
-ballangvel
-framelinvel
-frameangvel
* acceleration/force-dependent sensors:
-actuatorfrc
-jointactfrc
@@ -46,7 +48,7 @@
<!-- body 2 -->
<body name="body2" pos=".1 .1 .1">
<joint name="ballquat2" type="ball" pos="0.1 0.1 0.1"/>
<geom size="1"/>
<geom name="geom2" size="1"/>
</body>
<!-- body 3 -->
@@ -94,6 +96,8 @@
<jointvel name="jointvel0" joint="hinge0"/>
<jointvel name="jointvel0cutoff" joint="hinge0" cutoff="1e-3"/>
<actuatorfrc name="actuatorfrc0" actuator="motor0"/>
<framelinvel name="framelinvel3" objtype="site" objname="site3"/>
<framelinvel name="framelinvel3ref" objtype="site" objname="site3" reftype="body" refname="body2"/>
<actuatorpos name="actuatorpos0" actuator="motor0"/>
<actuatorvel name="actuatorvel0" actuator="motor0"/>
<ballquat name="ballquat2" joint="ballquat2"/>
@@ -106,6 +110,9 @@
<framepos name="framepos0" objtype="site" objname="site0"/>
<velocimeter name="velocimeter1" site="site1"/>
<gyro name="gyro1" site="site1"/>
<frameangvel name="frameangvel3" objtype="site" objname="site3"/>
<frameangvel name="frameangvel3ref" objtype="site" objname="site3" reftype="body" refname="body2"/>
<frameangvel name="frameangvel3cutoff" objtype="site" objname="site3" cutoff="1e-3"/>
<actuatorfrc name="actuatorfrc1" actuator="motor1"/>
<subtreecom name="subtreecom0" body="body0"/>
<camprojection site="frontorigin" camera="fixedcamera"/>
@@ -121,10 +128,15 @@
<ballquat name="ballquat3" joint="ballquat3"/>
<camprojection site="frontcenter" camera="fixedcamera"/>
<ballangvel name="ballangvel3" joint="ballquat3"/>
<frameangvel name="frameangvel0" objtype="site" objname="site0"/>
<frameangvel name="frameangvel0ref" objtype="site" objname="site0" reftype="site" refname="site3"/>
<frameyaxis name="frameyaxis1" objtype="site" objname="site1"/>
<framequat name="framequat0" objtype="site" objname="site0"/>
<rangefinder name="rangefinder1" site="site_rangefinder1"/>
<subtreecom name="subtreecom1" body="body1"/>
<framelinvel name="framelinvel0" objtype="site" objname="site0"/>
<framelinvel name="framelinvel0ref" objtype="site" objname="site0" reftype="geom" refname="geom2"/>
<framelinvel name="framelinvel0cutoff" objtype="site" objname="site0" cutoff="3e-4"/>
<clock/>
</sensor>
</mujoco>