diff --git a/mjx/mujoco/mjx/_src/support.py b/mjx/mujoco/mjx/_src/support.py index 7b17fcfa..421126f9 100644 --- a/mjx/mujoco/mjx/_src/support.py +++ b/mjx/mujoco/mjx/_src/support.py @@ -277,3 +277,43 @@ def name2id( } return names_map.get(name, -1) + + +def _decode_pyramid( + pyramid: jax.Array, mu: jax.Array, condim: int +) -> jax.Array: + """Converts pyramid representation to contact force.""" + force = jp.zeros(6, dtype=float) + if condim == 1: + return force.at[0].set(pyramid[0]) + + # force_normal = sum(pyramid0_i + pyramid1_i) + force = force.at[0].set(pyramid[0 : 2 * (condim - 1)].sum()) + + # force_tangent_i = (pyramid0_i - pyramid1_i) * mu_i + i = np.arange(0, condim) + force = force.at[i + 1].set((pyramid[2 * i] - pyramid[2 * i + 1]) * mu[i]) + + return force + + +def contact_force( + m: Model, d: Data, contact_id: int, to_world_frame: bool = False +) -> jax.Array: + """Extract 6D force:torque for one contact, in contact frame by default.""" + efc_address = d.contact.efc_address[contact_id] + condim = d.contact.dim[contact_id] + if m.opt.cone == mujoco.mjtCone.mjCONE_PYRAMIDAL: + force = _decode_pyramid( + d.efc_force[efc_address:], d.contact.friction[contact_id], condim + ) + elif m.opt.cone == mujoco.mjtCone.mjCONE_ELLIPTIC: + raise NotImplementedError('Elliptic cone force is not implemented yet.') + else: + raise ValueError(f'Unknown cone type: {m.opt.cone}') + + if to_world_frame: + force = force.reshape((-1, 3)) @ d.contact.frame[contact_id] + force = force.reshape(-1) + + return force * (efc_address >= 0) diff --git a/mjx/mujoco/mjx/_src/support_test.py b/mjx/mujoco/mjx/_src/support_test.py index 4aaebad4..ae539f15 100644 --- a/mjx/mujoco/mjx/_src/support_test.py +++ b/mjx/mujoco/mjx/_src/support_test.py @@ -157,6 +157,65 @@ class SupportTest(parameterized.TestCase): i = i if n is not None else -1 self.assertEqual(support.name2id(mx, obj, n), i) + _CONTACTS = """ + + + + + + + + + + + + + + + + + """ + + def test_contact_force(self): + m = mujoco.MjModel.from_xml_string(self._CONTACTS) + d = mujoco.MjData(m) + mujoco.mj_step(m, d) + assert ( + np.unique(d.contact.geom).shape[0] == 3 + ), 'This test assumes all capsule are in contact.' + mx = mjx.put_model(m) + dx = mjx.put_data(m, d) + mujoco.mj_step(m, d) + dx = mjx.step(mx, dx) + + # map MJX contacts to MJ ones + def _find(g): + val = (g == dx.contact.geom).sum(axis=1) + return np.where(val == 2)[0][0] + + contact_id_map = {i: _find(d.contact.geom[i]) for i in range(d.ncon)} + + for i in range(d.ncon): + result = np.zeros(6, dtype=float) + mujoco.mj_contactForce(m, d, i, result) + + j = contact_id_map[i] + force = jax.jit(support.contact_force, static_argnums=(2,))(mx, dx, j) + np.testing.assert_allclose(result, force, rtol=1e-5, atol=2) + + # test world conversion + force = jax.jit( + support.contact_force, + static_argnums=( + 2, + 3, + ), + )(mx, dx, j, True) + # back to contact frame + force = force.at[:3].set(dx.contact.frame[j] @ force[:3]) + force = force.at[3:].set(dx.contact.frame[j] @ force[3:]) + np.testing.assert_allclose(result, force, rtol=1e-5, atol=2) + if __name__ == '__main__': absltest.main() diff --git a/mjx/mujoco/mjx/test_data/convex.xml b/mjx/mujoco/mjx/test_data/convex.xml index 1006cb7a..382273d7 100644 --- a/mjx/mujoco/mjx/test_data/convex.xml +++ b/mjx/mujoco/mjx/test_data/convex.xml @@ -1,8 +1,4 @@ - - - -