From da9037e2b0a8848109f3737533c9720310fcc87d Mon Sep 17 00:00:00 2001 From: Baruch Tabanpour Date: Tue, 12 Aug 2025 02:22:53 -0700 Subject: [PATCH] Fix bug in MJX contact sensor. PiperOrigin-RevId: 794011199 Change-Id: I709a8bd9d07bec980aa60abd0346d946b037cbbd --- mjx/mujoco/mjx/_src/sensor.py | 13 +++++-------- mjx/mujoco/mjx/_src/sensor_test.py | 5 +---- 2 files changed, 6 insertions(+), 12 deletions(-) diff --git a/mjx/mujoco/mjx/_src/sensor.py b/mjx/mujoco/mjx/_src/sensor.py index 1ab91465..31f7f030 100644 --- a/mjx/mujoco/mjx/_src/sensor.py +++ b/mjx/mujoco/mjx/_src/sensor.py @@ -588,6 +588,7 @@ def sensor_acc(m: Model, d: Data) -> Data: size = nslotdata(dataspec) num = np.minimum(int(dim / size), ncon) + nsensor = idx_ds.sum() if objtype == ObjType.UNKNOWN and reftype == ObjType.UNKNOWN: # all contacts match @@ -601,11 +602,9 @@ def sensor_acc(m: Model, d: Data) -> Data: nfound = sum(is_contact) # if duplicate sensor - nsensor = idx_ds.sum() - cid = np.tile(cid, (nsensor,)) - match = np.tile(match[:num], (nsensor,)) - nfound = np.tile(nfound, (nsensor,)) - flip = np.ones((cid.size, 3)) + cid = jp.tile(cid, (nsensor,)) + nfound = jp.tile(nfound, (nsensor,)) + flip = jp.ones((cid.size, 3)) elif objtype == ObjType.GEOM or reftype == ObjType.GEOM: sensorid1 = objid[idx_ds] sensorid2 = refid[idx_ds] @@ -651,8 +650,6 @@ def sensor_acc(m: Model, d: Data) -> Data: # number of contacts per sensor nfound = (match * is_contact[None, :]).sum(axis=1) - match = match[:, :num].reshape(-1) - # TODO(taylorhowell): matching criteria: body, subtree else: @@ -683,7 +680,7 @@ def sensor_acc(m: Model, d: Data) -> Data: if dataspec & (1 << 6): # tangent slot.append(flip[:, 2, None] * d._impl.contact.frame[cid, 1]) - found = is_contact[cid] & match + found = jp.tile(jp.arange(num), nsensor) < jp.repeat(nfound, num) sensors.append((found[:, None] * jp.hstack(slot)).reshape(-1)) adrs.append( (adr[idx_ds][:, None] + np.arange(num * size)[None]).reshape(-1) diff --git a/mjx/mujoco/mjx/_src/sensor_test.py b/mjx/mujoco/mjx/_src/sensor_test.py index 60ae5222..2843473c 100644 --- a/mjx/mujoco/mjx/_src/sensor_test.py +++ b/mjx/mujoco/mjx/_src/sensor_test.py @@ -114,12 +114,11 @@ class SensorTest(parameterized.TestCase): datas = list(datas) contact_sensors = '' - for num in [1, 2, 3, 4, 5]: + for num in [1, 2, 4, 5]: for data in datas: data = ' '.join(data) for reduce in ['mindist', 'maxforce']: for match in [ - '', '', 'geom1="plane"', 'geom1="geom1"', @@ -130,9 +129,7 @@ class SensorTest(parameterized.TestCase): 'geom1="plane" geom2="geom1"', 'geom1="geom1" geom2="plane"', 'geom1="plane" geom2="sphere2"', - 'geom1="sphere2" geom2="plane"', 'geom1="geom1" geom2="sphere2"', - 'geom1="sphere2" geom2="geom1"', ]: contact_sensors += ( f'