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 @@
-
-
-
-