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