diff --git a/doc/changelog.rst b/doc/changelog.rst index 9152b8a7..4ddf86f1 100644 --- a/doc/changelog.rst +++ b/doc/changelog.rst @@ -17,8 +17,9 @@ General MJX ^^^ 5. Added ``efc_pos`` to ``mjx.Data`` (:github:issue:`1388`). -6. Added position-dependent sensors: ``MAGNETOMETER``, ``RANGEFINDER``, ``JOINTPOS``, ``ACTUATORPOS``, ``BALLQUAT``, - ``FRAMEPOS``, ``FRAMEXAXIS``, ``FRAMEYAXIS``, ``FRAMEZAXIS``, ``SUBTREECOM``, ``CLOCK``. +6. Added position-dependent sensors: ``MAGNETOMETER``, ``CAMPROJECTION``, ``RANGEFINDER``, ``JOINTPOS``, + ``ACTUATORPOS``, ``BALLQUAT``, ``FRAMEPOS``, ``FRAMEXAXIS``, ``FRAMEYAXIS``, ``FRAMEZAXIS``, ``SUBTREECOM``, + ``CLOCK``. 7. Changed default policy to avoid placing unused (MuJoCo-only) arrays on device. 8. Added ``device`` parameter to ``mjx.make_data`` to bring it to parity with ``mjx.put_model`` and ``mjx.put_data``. diff --git a/mjx/mujoco/mjx/_src/sensor.py b/mjx/mujoco/mjx/_src/sensor.py index afe3741d..f8792764 100644 --- a/mjx/mujoco/mjx/_src/sensor.py +++ b/mjx/mujoco/mjx/_src/sensor.py @@ -72,6 +72,66 @@ def sensor_pos(m: Model, d: Data) -> Data: d.site_xmat[objid] ).reshape(-1) adr = (adr[:, None] + np.arange(3)[None]).reshape(-1) + elif sensor_type == SensorType.CAMPROJECTION: + + @jax.vmap + def _cam_project( + target_xpos, xpos, xmat, res, fovy, intrinsic, sensorsize, focal_flag + ): + translation = jp.eye(4).at[0:3, 3].set(-xpos) + rotation = jp.eye(4).at[:3, :3].set(xmat.T) + + # focal transformation matrix (3 x 4) + f = 0.5 / jp.tan(fovy * jp.pi / 360.0) * res[1] + fx, fy = jp.where( + focal_flag, + intrinsic[:2] / (sensorsize[:2] + mujoco.mjMINVAL) * res[:2], + f, + ) # add mjMINVAL to denominator to prevent divide by zero warning + + focal = jp.array([[-fx, 0, 0, 0], [0, fy, 0, 0], [0, 0, 1.0, 0]]) + + # image matrix (3 x 3) + image = jp.eye(3).at[:2, 2].set(res[0:2] / 2.0) + + # projection matrix (3 x 4): product of all 4 matrices + proj = image @ focal @ rotation @ translation + + # projection matrix multiplies homogenous [x, y, z, 1] vectors + pos_hom = jp.append(target_xpos, 1.0) + + # project world coordinates into pixel space, see: + # https://en.wikipedia.org/wiki/3D_projection#Mathematical_formula + pixel_coord_hom = proj @ pos_hom + + # avoid dividing by tiny numbers + denom = pixel_coord_hom[2] + denom = jp.where( + jp.abs(denom) < mujoco.mjMINVAL, + jp.clip(denom, -mujoco.mjMINVAL, mujoco.mjMINVAL), + denom, + ) + + # compute projection + sensor = pixel_coord_hom / denom + + return sensor[:2] + + refid = m.sensor_refid[idx] + sensorsize = m.cam_sensorsize[refid] + intrinsic = m.cam_intrinsic[refid] + fovy = m.cam_fovy[refid] + res = m.cam_resolution[refid] + focal_flag = np.logical_and(sensorsize[:, 0] != 0, sensorsize[:, 1] != 0) + + target_xpos = d.site_xpos[objid] + xpos = d.cam_xpos[refid] + xmat = d.cam_xmat[refid] + + 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] for sid in set(site_bodyid): diff --git a/mjx/mujoco/mjx/_src/types.py b/mjx/mujoco/mjx/_src/types.py index b92db752..de53a7a1 100644 --- a/mjx/mujoco/mjx/_src/types.py +++ b/mjx/mujoco/mjx/_src/types.py @@ -293,6 +293,7 @@ class SensorType(enum.IntEnum): Members: MAGNETOMETER: magnetometer + CAMPROJECTION: camera projection RANGEFINDER: rangefinder JOINTPOS: joint position ACTUATORPOS: actuator position @@ -305,6 +306,7 @@ class SensorType(enum.IntEnum): CLOCK: simulation time """ MAGNETOMETER = mujoco.mjtSensor.mjSENS_MAGNETOMETER + CAMPROJECTION = mujoco.mjtSensor.mjSENS_CAMPROJECTION RANGEFINDER = mujoco.mjtSensor.mjSENS_RANGEFINDER JOINTPOS = mujoco.mjtSensor.mjSENS_JOINTPOS ACTUATORPOS = mujoco.mjtSensor.mjSENS_ACTUATORPOS diff --git a/mjx/mujoco/mjx/test_data/sensor.xml b/mjx/mujoco/mjx/test_data/sensor.xml index 80763173..72b21e3f 100644 --- a/mjx/mujoco/mjx/test_data/sensor.xml +++ b/mjx/mujoco/mjx/test_data/sensor.xml @@ -2,6 +2,7 @@ * position-dependent sensors: -magnetometer +-camprojection -rangefinder -jointpos -actuatorpos @@ -46,6 +47,19 @@ + + + + + + + + + + + + @@ -60,7 +74,6 @@ - @@ -71,6 +84,7 @@ + @@ -78,11 +92,10 @@ + - -