From 1c424644dd68875dba95e973e171b358c4f0f5a4 Mon Sep 17 00:00:00 2001
From: Taylor Howell
Date: Mon, 28 Oct 2024 03:38:52 -0700
Subject: [PATCH] Fix MJX touch sensor.
PiperOrigin-RevId: 690544118
Change-Id: Ic909c6f6ce0e31237cbe317fbbfbb1fd267bd462
---
mjx/mujoco/mjx/_src/sensor.py | 21 ++++++++---------
mjx/mujoco/mjx/_src/sensor_test.py | 3 ---
mjx/mujoco/mjx/test_data/sensor/sensor.xml | 27 ++++++++++++++--------
3 files changed, 28 insertions(+), 23 deletions(-)
diff --git a/mjx/mujoco/mjx/_src/sensor.py b/mjx/mujoco/mjx/_src/sensor.py
index e514bdd2..5d040052 100644
--- a/mjx/mujoco/mjx/_src/sensor.py
+++ b/mjx/mujoco/mjx/_src/sensor.py
@@ -461,7 +461,7 @@ def sensor_acc(m: Model, d: Data) -> Data:
force, condim_id = support.contact_force_dim(m, d, dim)
forces.append(force)
condim_ids.append(condim_id)
- forces = jp.concatenate(forces)[jp.concatenate(condim_ids)]
+ forces = jp.concatenate(forces)[np.argsort(np.concatenate(condim_ids))]
# get bodies of contact geoms
conbody = jp.array(m.geom_bodyid)[d.contact.geom]
@@ -483,14 +483,14 @@ def sensor_acc(m: Model, d: Data) -> Data:
conray = jp.where(conbody1[..., None], -conray, conray)
# compute distance, mapping over sites and contacts
- def _distance(
- site_size, site_xpos, site_xmat, site_type, contact_pos, conray
- ):
- return jax.vmap(
- lambda site_size, site_xpos, site_xmat, conray: jax.vmap(
- lambda pnt, vec: ray.ray_geom(site_size, pnt, vec, site_type)
- )((contact_pos - site_xpos) @ site_xmat, conray @ site_xmat)
- )(site_size, site_xpos, site_xmat, conray)
+ def _distance(site_size, site_xpos, site_xmat, site_type, pos, conray):
+ def dist(size, xpos, xmat, conray):
+ pnt = (pos - xpos) @ xmat
+ vec = conray @ xmat
+ ray_geom_ = lambda pnt, vec: ray.ray_geom(size, pnt, vec, site_type)
+ return jax.vmap(ray_geom_)(pnt, vec)
+
+ return jax.vmap(dist)(site_size, site_xpos, site_xmat, conray)
dist = []
dist_id = []
@@ -506,8 +506,7 @@ def sensor_acc(m: Model, d: Data) -> Data:
)
dist.append(jp.where(jp.isinf(dist_site), 0, dist_site))
dist_id.append(dist_id_site)
-
- dist = jp.vstack(dist)[np.concatenate(dist_id)]
+ dist = jp.vstack(dist)[np.argsort(np.concatenate(dist_id))]
# accumulate normal forces for each site
sensor = jp.dot((dist > 0) & contacts, forces[:, 0])
diff --git a/mjx/mujoco/mjx/_src/sensor_test.py b/mjx/mujoco/mjx/_src/sensor_test.py
index 9eedac13..19d4ef46 100644
--- a/mjx/mujoco/mjx/_src/sensor_test.py
+++ b/mjx/mujoco/mjx/_src/sensor_test.py
@@ -101,17 +101,14 @@ class SensorTest(parameterized.TestCase):
-
-
-
""")
diff --git a/mjx/mujoco/mjx/test_data/sensor/sensor.xml b/mjx/mujoco/mjx/test_data/sensor/sensor.xml
index c8257831..71532743 100644
--- a/mjx/mujoco/mjx/test_data/sensor/sensor.xml
+++ b/mjx/mujoco/mjx/test_data/sensor/sensor.xml
@@ -102,15 +102,22 @@
-
-
+
+
+
+
-
-
+
+
+
+
+
+
+
+
+
+
@@ -141,7 +148,8 @@
-
+
+
@@ -152,7 +160,8 @@
-
+
+