diff --git a/mjx/mujoco/mjx/__init__.py b/mjx/mujoco/mjx/__init__.py
index d272dd13..fef95c32 100644
--- a/mjx/mujoco/mjx/__init__.py
+++ b/mjx/mujoco/mjx/__init__.py
@@ -43,6 +43,7 @@ from mujoco.mjx._src.smooth import crb
from mujoco.mjx._src.smooth import factor_m
from mujoco.mjx._src.smooth import kinematics
from mujoco.mjx._src.smooth import rne
+from mujoco.mjx._src.smooth import rne_postconstraint
from mujoco.mjx._src.smooth import subtree_vel
from mujoco.mjx._src.smooth import tendon
from mujoco.mjx._src.smooth import transmission
diff --git a/mjx/mujoco/mjx/_src/sensor_test.py b/mjx/mujoco/mjx/_src/sensor_test.py
index 0ff11a86..e2b7bd35 100644
--- a/mjx/mujoco/mjx/_src/sensor_test.py
+++ b/mjx/mujoco/mjx/_src/sensor_test.py
@@ -40,7 +40,7 @@ def _assert_attr_eq(a, b, attr):
class SensorTest(parameterized.TestCase):
- @parameterized.parameters('no_sensor.xml', 'sensor.xml')
+ @parameterized.parameters('sensor/model.xml', 'sensor/sensor.xml')
def test_sensor(self, filename):
"""Tests MJX sensor functions match MuJoCo sensor functions."""
m = test_util.load_test_file(filename)
@@ -65,7 +65,7 @@ class SensorTest(parameterized.TestCase):
def test_disable_sensor(self):
"""Tests disabling sensor."""
- m = test_util.load_test_file('sensor.xml')
+ m = test_util.load_test_file('sensor/sensor.xml')
# disable sensors
m.opt.disableflags = m.opt.disableflags | mjx.DisableBit.SENSOR
d = mujoco.MjData(m)
@@ -84,7 +84,7 @@ class SensorTest(parameterized.TestCase):
def test_unsupported_sensor(self):
"""Tests MJX sensor functions do not break for unsupported sensors."""
- m = test_util.load_test_file('unsupported_sensor.xml')
+ m = test_util.load_test_file('sensor/unsupported.xml')
mx = mjx.put_model(m)
dx = jax.jit(mjx.forward)(mx, mjx.put_data(m, mujoco.MjData(m)))
_assert_eq(np.zeros(m.nsensordata), dx.sensordata, 'sensordata')
diff --git a/mjx/mujoco/mjx/_src/smooth.py b/mjx/mujoco/mjx/_src/smooth.py
index 96918c6a..6d503406 100644
--- a/mjx/mujoco/mjx/_src/smooth.py
+++ b/mjx/mujoco/mjx/_src/smooth.py
@@ -558,6 +558,129 @@ def rne(m: Model, d: Data) -> Data:
return d
+def rne_postconstraint(m: Model, d: Data) -> Data:
+ """RNE with complete data: compute cacc, cfrc_ext, cfrc_int."""
+
+ def _transform_force(frc, offset):
+ force, torque = jp.split(frc, 2)
+ torque -= jp.cross(offset, force)
+ # spatial motion vector layout is flipped: (torque, force)
+ return jp.concatenate([torque, force])
+
+ # cfrc_ext = perturb
+ cfrc_ext = jp.vstack([
+ jp.zeros((1, 6)), # world body
+ jax.vmap(_transform_force)(
+ d.xfrc_applied[1:], d.subtree_com[m.body_rootid][1:] - d.xipos[1:]
+ ),
+ ])
+
+ # cfrc_ext += contacts
+
+ # compute contact forces for each condim
+ forces = []
+ condim_idx = []
+ for dim in set(d.contact.dim):
+ force, idx = support.contact_force_dim(m, d, dim)
+ forces.append(force)
+ condim_idx.append(idx)
+
+ # update cfrc_ext with contact forces
+ if forces:
+
+ @jax.vmap
+ def _contact_force_to_cfrc_ext(force, pos, frame, id1, id2, com1, com2):
+ # force: contact to world frame
+ force = force.reshape((-1, 3)) @ frame
+ force = force.reshape(-1)
+
+ # contact force on bodies
+ cfrc_com1 = _transform_force(force, com1 - pos)
+ cfrc_com2 = _transform_force(force, com2 - pos)
+
+ # mask
+ mask1 = id1 != 0
+ mask2 = id2 != 0
+
+ return jp.vstack([-1 * cfrc_com1 * mask1, cfrc_com2 * mask2]), jp.array(
+ [id1, id2]
+ )
+
+ condim_idx = jp.concatenate(condim_idx)
+ frame = d.contact.frame[condim_idx]
+ pos = d.contact.pos[condim_idx]
+ id1 = jp.array(m.geom_bodyid)[d.contact.geom[condim_idx, 0]]
+ id2 = jp.array(m.geom_bodyid)[d.contact.geom[condim_idx, 1]]
+ com1 = d.subtree_com[jp.array(m.body_rootid)][id1]
+ com2 = d.subtree_com[jp.array(m.body_rootid)][id2]
+
+ cfrc_contact, cfrc_idx = _contact_force_to_cfrc_ext(
+ jp.concatenate(forces), pos, frame, id1, id2, com1, com2
+ )
+
+ cfrc_ext = cfrc_ext.at[cfrc_idx.reshape(-1)].add(
+ cfrc_contact.reshape((-1, 6))
+ )
+
+ # TODO(taylorhowell): connect and weld constraints
+
+ # forward pass over bodies: compute cacc, cfrc_int
+ def _forward(carry, cfrc_ext, cinert, cvel, body_dofadr, body_dofnum):
+ if carry is None:
+ if m.opt.disableflags & DisableBit.GRAVITY:
+ cacc0 = jp.zeros(6)
+ else:
+ cacc0 = jp.concatenate((jp.zeros(3), -m.opt.gravity))
+ return cacc0, jp.zeros(6)
+ else:
+ cacc_parent, _ = carry
+
+ # create dof mask
+ indices = jp.arange(m.nv)
+ mask = jp.logical_and(
+ indices >= body_dofadr, indices < body_dofadr + body_dofnum
+ )
+
+ # cacc = cacc_parent + cdofdot * qvel + cdof * qacc
+ cacc_vel = d.cdof_dot.T @ (mask * d.qvel)
+ cacc_acc = d.cdof.T @ (mask * d.qacc)
+ cacc = cacc_parent + cacc_vel + cacc_acc
+
+ # cfrc_body = cinert * cacc + cvel x (cinert * cvel)
+ cfrc_body = math.inert_mul(cinert, cacc)
+ cfrc_corr = math.inert_mul(cinert, cvel)
+ cfrc = math.motion_cross_force(cvel, cfrc_corr)
+ cfrc_body = cfrc_body + cfrc
+ cfrc_int = cfrc_body - cfrc_ext
+
+ return cacc, cfrc_int
+
+ cacc, cfrc_int = scan.body_tree(
+ m,
+ _forward,
+ 'bbbbb',
+ 'bb',
+ cfrc_ext,
+ d.cinert,
+ d.cvel,
+ jp.array(m.body_dofadr),
+ jp.array(m.body_dofnum),
+ )
+
+ # backward pass over bodies: accumulate cfrc_int from children
+ cfrc_int = scan.body_tree(
+ m,
+ lambda c, p: p + c if c is not None else p, # add child to parent
+ 'b',
+ 'b',
+ cfrc_int,
+ reverse=True,
+ )
+
+ # update data
+ return d.replace(cacc=cacc, cfrc_int=cfrc_int, cfrc_ext=cfrc_ext)
+
+
def tendon(m: Model, d: Data) -> Data:
"""Computes tendon lengths and moments."""
if not m.ntendon:
diff --git a/mjx/mujoco/mjx/_src/smooth_test.py b/mjx/mujoco/mjx/_src/smooth_test.py
index c8e5f162..c94347ce 100644
--- a/mjx/mujoco/mjx/_src/smooth_test.py
+++ b/mjx/mujoco/mjx/_src/smooth_test.py
@@ -191,6 +191,44 @@ class SmoothTest(absltest.TestCase):
_assert_attr_eq(d, dx, 'subtree_linvel')
_assert_attr_eq(d, dx, 'subtree_angmom')
+ def test_rnepostconstraint(self):
+ """Tests MJX rne_postconstraint function to match MuJoCo mj_rnePostConstraint."""
+
+ m = mujoco.MjModel.from_xml_string("""
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+ """)
+ d = mujoco.MjData(m)
+ mujoco.mj_resetDataKeyframe(m, d, 0)
+ # apply external forces
+ d.xfrc_applied = 0.001 * np.ones(d.xfrc_applied.shape)
+ mujoco.mj_step(m, d, 2)
+ mujoco.mj_forward(m, d)
+ mx = mjx.put_model(m)
+ dx = mjx.put_data(m, d)
+
+ # rne postconstraint
+ mujoco.mj_rnePostConstraint(m, d)
+ dx = jax.jit(mjx.rne_postconstraint)(mx, dx)
+
+ _assert_eq(d.cacc, dx.cacc, 'cacc')
+ _assert_eq(d.cfrc_ext, dx.cfrc_ext, 'cfrc_ext')
+ _assert_eq(d.cfrc_int, dx.cfrc_int, 'cfrc_int')
+
if __name__ == '__main__':
absltest.main()
diff --git a/mjx/mujoco/mjx/_src/support.py b/mjx/mujoco/mjx/_src/support.py
index 0cd31636..03e510a0 100644
--- a/mjx/mujoco/mjx/_src/support.py
+++ b/mjx/mujoco/mjx/_src/support.py
@@ -319,3 +319,28 @@ def contact_force(
force = force.reshape(-1)
return force * (efc_address >= 0)
+
+
+def contact_force_dim(
+ m: Model, d: Data, dim: int
+) -> Tuple[jax.Array, np.ndarray]:
+ """Extract 6D force:torque for contacts with dimension dim."""
+ # valid contact and condim indices
+ idx_dim = (d.contact.efc_address >= 0) & (d.contact.dim == dim)
+
+ # contact force from efc
+ if m.opt.cone == mujoco.mjtCone.mjCONE_PYRAMIDAL:
+ efc_address = (
+ d.contact.efc_address[idx_dim, None]
+ + np.arange(np.where(dim == 1, 1, 2 * (dim - 1)))[None]
+ )
+ efc_force = d.efc_force[efc_address]
+ force = jax.vmap(_decode_pyramid, in_axes=(0, 0, None))(
+ efc_force, d.contact.friction[idx_dim], dim
+ )
+ return force, np.where(idx_dim)[0]
+ elif m.opt.cone == mujoco.mjtCone.mjCONE_ELLIPTIC:
+ # TODO(taylorhowell): add support for elliptic cone
+ raise NotImplementedError('Elliptic cone force is not implemented yet.')
+ else:
+ raise ValueError(f'Unknown cone type: {m.opt.cone}.')
diff --git a/mjx/mujoco/mjx/test_data/no_sensor.xml b/mjx/mujoco/mjx/test_data/no_sensor.xml
deleted file mode 100644
index 5ccc660b..00000000
--- a/mjx/mujoco/mjx/test_data/no_sensor.xml
+++ /dev/null
@@ -1,9 +0,0 @@
-
-
-
-
-
-
-
-
-
diff --git a/mjx/mujoco/mjx/test_data/sensor/model.xml b/mjx/mujoco/mjx/test_data/sensor/model.xml
new file mode 100644
index 00000000..18726d46
--- /dev/null
+++ b/mjx/mujoco/mjx/test_data/sensor/model.xml
@@ -0,0 +1,58 @@
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
diff --git a/mjx/mujoco/mjx/test_data/sensor.xml b/mjx/mujoco/mjx/test_data/sensor/sensor.xml
similarity index 100%
rename from mjx/mujoco/mjx/test_data/sensor.xml
rename to mjx/mujoco/mjx/test_data/sensor/sensor.xml
diff --git a/mjx/mujoco/mjx/test_data/sensor/unsupported.xml b/mjx/mujoco/mjx/test_data/sensor/unsupported.xml
new file mode 100644
index 00000000..30ba0e70
--- /dev/null
+++ b/mjx/mujoco/mjx/test_data/sensor/unsupported.xml
@@ -0,0 +1,65 @@
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
diff --git a/mjx/mujoco/mjx/test_data/unsupported_sensor.xml b/mjx/mujoco/mjx/test_data/unsupported_sensor.xml
deleted file mode 100644
index 4a164bd9..00000000
--- a/mjx/mujoco/mjx/test_data/unsupported_sensor.xml
+++ /dev/null
@@ -1,19 +0,0 @@
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-