From e9709900b4b395486253fb04e657c1949cc2c4e9 Mon Sep 17 00:00:00 2001 From: Baruch Tabanpour Date: Wed, 24 Apr 2024 11:18:35 -0700 Subject: [PATCH] Add gravcomp. PiperOrigin-RevId: 627791303 Change-Id: I87c9254967eb492d0c38153ff41d7458ffb69136 --- doc/changelog.rst | 9 +- mjx/mujoco/mjx/_src/forward.py | 9 +- mjx/mujoco/mjx/_src/io.py | 4 +- mjx/mujoco/mjx/_src/io_test.py | 12 -- mjx/mujoco/mjx/_src/passive.py | 162 ++++++++++-------- mjx/mujoco/mjx/_src/passive_test.py | 6 +- mjx/mujoco/mjx/_src/test_util.py | 12 +- mjx/mujoco/mjx/_src/types.py | 11 +- .../mjx/integration_test/smooth_test.py | 8 +- mjx/mujoco/mjx/test_data/pendula.xml | 17 ++ 10 files changed, 152 insertions(+), 98 deletions(-) diff --git a/doc/changelog.rst b/doc/changelog.rst index d705b30d..849d806e 100644 --- a/doc/changelog.rst +++ b/doc/changelog.rst @@ -40,14 +40,15 @@ MJX 12. Added cylinder collisions using SDFs. 13. Added support for all :ref:`condim `: 1, 3, 4, 6. 14. Add support functions for ``id2name`` and ``name2id``, MJX versions of :ref:`mj_id2name` and :ref:`mj_name2id`. +15. Added support for :ref:`gravcomp` and :ref:`actuatorgravcomp`. Bug fixes ^^^^^^^^^ -15. Defaults of lights were not being saved, now fixed. -16. Prevent overwriting of frame names by body names when saving an XML. Bug introduced in 3.1.4. -17. Fixed bug in Python binding of :ref:`mj_saveModel`: ``buffer`` argument was documented as optional but was actually +16. Defaults of lights were not being saved, now fixed. +17. Prevent overwriting of frame names by body names when saving an XML. Bug introduced in 3.1.4. +18. Fixed bug in Python binding of :ref:`mj_saveModel`: ``buffer`` argument was documented as optional but was actually not optional. -18. Fixed bug that prevented memory allocations larger than 2.15 GB. +19. Fixed bug that prevented memory allocations larger than 2.15 GB. Version 3.1.4 (April 10th, 2024) diff --git a/mjx/mujoco/mjx/_src/forward.py b/mjx/mujoco/mjx/_src/forward.py index 905e463a..e2838d00 100644 --- a/mjx/mujoco/mjx/_src/forward.py +++ b/mjx/mujoco/mjx/_src/forward.py @@ -174,16 +174,17 @@ def fwd_actuation(m: Model, d: Data) -> Data: qfrc_actuator = d.actuator_moment.T @ force + if m.ngravcomp: + # actuator-level gravity compensation, skip if added as passive force + qfrc_actuator += d.qfrc_gravcomp * m.jnt_actgravcomp[m.dof_jntid] + # clamp qfrc_actuator actfrcrange = jp.where( m.jnt_actfrclimited[:, None], m.jnt_actfrcrange, jp.array([-jp.inf, jp.inf]), ) - ids = sum( - ([i] * JointType(j).dof_width() for i, j in enumerate(m.jnt_type)), [] - ) - actfrcrange = jp.take(actfrcrange, jp.array(ids), axis=0) + actfrcrange = actfrcrange[m.dof_jntid] qfrc_actuator = jp.clip(qfrc_actuator, actfrcrange[:, 0], actfrcrange[:, 1]) d = d.replace(act_dot=act_dot, qfrc_actuator=qfrc_actuator) diff --git a/mjx/mujoco/mjx/_src/io.py b/mjx/mujoco/mjx/_src/io.py index 1870e2dd..594ca15f 100644 --- a/mjx/mujoco/mjx/_src/io.py +++ b/mjx/mujoco/mjx/_src/io.py @@ -68,9 +68,6 @@ def put_model(m: mujoco.MjModel, device=None) -> types.Model: if m.ntendon: raise NotImplementedError('tendons are not supported') - if m.body_gravcomp.any(): - raise NotImplementedError('gravcomp is not supported') - for g1, g2, ip in collision_driver.geom_pairs(m): t1, t2 = m.geom_type[[g1, g2]] # check collision function exists for type pair @@ -201,6 +198,7 @@ def make_data(m: Union[types.Model, mujoco.MjModel]) -> types.Data: cvel=zero_nbody_6, cdof_dot=zero_nv_6, qfrc_bias=zero_nv, + qfrc_gravcomp=zero_nv, qfrc_passive=zero_nv, efc_aref=zero_nefc, qfrc_actuator=zero_nv, diff --git a/mjx/mujoco/mjx/_src/io_test.py b/mjx/mujoco/mjx/_src/io_test.py index 72f8b5c2..525db2bb 100644 --- a/mjx/mujoco/mjx/_src/io_test.py +++ b/mjx/mujoco/mjx/_src/io_test.py @@ -173,18 +173,6 @@ class ModelIOTest(parameterized.TestCase): """)) - def test_gravcomp_not_implemented(self): - with self.assertRaises(NotImplementedError): - mjx.put_model(mujoco.MjModel.from_xml_string(""" - - - - - - - - """)) - def test_cylinder_not_implemented(self): with self.assertRaises(NotImplementedError): mjx.put_model(mujoco.MjModel.from_xml_string(""" diff --git a/mjx/mujoco/mjx/_src/passive.py b/mjx/mujoco/mjx/_src/passive.py index 268de126..c2428566 100644 --- a/mjx/mujoco/mjx/_src/passive.py +++ b/mjx/mujoco/mjx/_src/passive.py @@ -29,6 +29,100 @@ from mujoco.mjx._src.types import Model # pylint: enable=g-importing-member +def _spring_damper(m: Model, d: Data) -> jax.Array: + """Applies joint level spring and damping forces.""" + def fn(jnt_typs, stiffness, qpos_spring, qpos): + qpos_i = 0 + qfrcs = [] + for i in range(len(jnt_typs)): + jnt_typ = JointType(jnt_typs[i]) + q = qpos[qpos_i : qpos_i + jnt_typ.qpos_width()] + qs = qpos_spring[qpos_i : qpos_i + jnt_typ.qpos_width()] + qfrc = jp.zeros(jnt_typ.dof_width()) + if jnt_typ == JointType.FREE: + qfrc = qfrc.at[:3].set(-stiffness[i] * (q[:3] - qs[:3])) + qfrc = qfrc.at[3:6].set(-stiffness[i] * math.quat_sub(q[3:7], qs[3:7])) + elif jnt_typ == JointType.BALL: + qfrc = -stiffness[i] * math.quat_sub(q, qs) + elif jnt_typ in ( + JointType.SLIDE, + JointType.HINGE, + ): + qfrc = -stiffness[i] * (q - qs) + else: + raise RuntimeError(f'unrecognized joint type: {jnt_typ}') + qfrcs.append(qfrc) + qpos_i += jnt_typ.qpos_width() + return jp.concatenate(qfrcs) + + # dof-level springs + qfrc = scan.flat( + m, + fn, + 'jjqq', + 'v', + m.jnt_type, + m.jnt_stiffness, + m.qpos_spring, + d.qpos, + ) + + # dof-level dampers + qfrc -= m.dof_damping * d.qvel + + return qfrc + + +def _gravcomp(m: Model, d: Data) -> jax.Array: + """Applies body-level gravity compensation.""" + force = -m.opt.gravity * (m.body_mass * m.body_gravcomp)[:, None] + + apply_f = lambda f, pos, body_id: support.jac(m, d, pos, body_id)[0] @ f + qfrc = jax.vmap(apply_f)(force, d.xipos, jp.arange(m.nbody)).sum(axis=0) + + return qfrc + + +def _fluid(m: Model, d: Data) -> jax.Array: + """Applies body-level viscosity, lift and drag.""" + force, torque = jax.vmap( + _inertia_box_fluid_model, in_axes=(None, 0, 0, 0, 0, 0, 0) + )( + m, + m.body_inertia, + m.body_mass, + d.subtree_com[jp.array(m.body_rootid)], + d.xipos, + d.ximat, + d.cvel, + ) + qfrc = jax.vmap(support.apply_ft, in_axes=(None, None, 0, 0, 0, 0))( + m, d, force, torque, d.xipos, jp.arange(m.nbody) + ) + + return jp.sum(qfrc, axis=0) + + +def passive(m: Model, d: Data) -> Data: + """Adds all passive forces.""" + if m.opt.disableflags & DisableBit.PASSIVE: + return d.replace(qfrc_passive=jp.zeros(m.nv), qfrc_gravcomp=jp.zeros(m.nv)) + + qfrc_passive = _spring_damper(m, d) + qfrc_gravcomp = jp.zeros(m.nv) + + if m.ngravcomp and not m.opt.disableflags & DisableBit.GRAVITY: + qfrc_gravcomp = _gravcomp(m, d) + # add gravcomp unless added via actuators + qfrc_passive += qfrc_gravcomp * (1 - m.jnt_actgravcomp[m.dof_jntid]) + + if m.opt.has_fluid_params: + qfrc_passive += _fluid(m, d) + + d = d.replace(qfrc_passive=qfrc_passive, qfrc_gravcomp=qfrc_gravcomp) + return d + + def _inertia_box_fluid_model( m: Model, inertia: jax.Array, @@ -71,71 +165,3 @@ def _inertia_box_fluid_model( force, torque = ximat @ lfrc_vel, ximat @ lfrc_ang return force, torque - - -def passive(m: Model, d: Data) -> Data: - """Adds all passive forces.""" - if m.opt.disableflags & DisableBit.PASSIVE: - return d.replace(qfrc_passive=jp.zeros(m.nv)) - - # joint-level springs - def fn(jnt_typs, stiffness, qpos_spring, qpos): - qpos_i = 0 - qfrcs = [] - for i in range(len(jnt_typs)): - jnt_typ = JointType(jnt_typs[i]) - q = qpos[qpos_i : qpos_i + jnt_typ.qpos_width()] - qs = qpos_spring[qpos_i : qpos_i + jnt_typ.qpos_width()] - qfrc = jp.zeros(jnt_typ.dof_width()) - if jnt_typ == JointType.FREE: - qfrc = qfrc.at[:3].set(-stiffness[i] * (q[:3] - qs[:3])) - qfrc = qfrc.at[3:6].set(-stiffness[i] * math.quat_sub(q[3:7], qs[3:7])) - elif jnt_typ == JointType.BALL: - qfrc = -stiffness[i] * math.quat_sub(q, qs) - elif jnt_typ in ( - JointType.SLIDE, - JointType.HINGE, - ): - qfrc = -stiffness[i] * (q - qs) - else: - raise RuntimeError(f'unrecognized joint type: {jnt_typ}') - qfrcs.append(qfrc) - qpos_i += jnt_typ.qpos_width() - return jp.concatenate(qfrcs) - - qfrc_passive = scan.flat( - m, - fn, - 'jjqq', - 'v', - m.jnt_type, - m.jnt_stiffness, - m.qpos_spring, - d.qpos, - ) - - # dof-level dampers - qfrc_passive -= m.dof_damping * d.qvel - - # TODO(robotics-simulation): body-level gravity compensation - - # body-level viscosity, lift and drag - if m.opt.has_fluid_params: - force, torque = jax.vmap( - _inertia_box_fluid_model, in_axes=(None, 0, 0, 0, 0, 0, 0) - )( - m, - m.body_inertia, - m.body_mass, - d.subtree_com[jp.array(m.body_rootid)], - d.xipos, - d.ximat, - d.cvel, - ) - qfrc_target = jax.vmap(support.apply_ft, in_axes=(None, None, 0, 0, 0, 0))( - m, d, force, torque, d.xipos, jp.arange(m.nbody) - ) - qfrc_passive += jp.sum(qfrc_target, axis=0) - - d = d.replace(qfrc_passive=qfrc_passive) - return d diff --git a/mjx/mujoco/mjx/_src/passive_test.py b/mjx/mujoco/mjx/_src/passive_test.py index 6ff5a596..910ce4d3 100644 --- a/mjx/mujoco/mjx/_src/passive_test.py +++ b/mjx/mujoco/mjx/_src/passive_test.py @@ -42,13 +42,14 @@ class PassiveTest(absltest.TestCase): m = test_util.load_test_file('pendula.xml') d = mujoco.MjData(m) # give the system a little kick to ensure we have non-identity rotations - d.ctrl = np.array([0.1, -0.1, 0.2, 0.3, -0.4]) + d.ctrl = np.array([0.1, -0.1, 0.2, 0.3, -0.4, 0.5, -0.6, 0.1]) mujoco.mj_step(m, d, 10) # let dynamics get state significantly non-zero mujoco.mj_forward(m, d) mx = mjx.put_model(m) dx = jax.jit(mjx.passive)(mx, mjx.put_data(m, d)) _assert_attr_eq(d, dx, 'qfrc_passive') + _assert_attr_eq(d, dx, 'qfrc_gravcomp') # test with fluid forces m.opt.density = 0.01 @@ -56,18 +57,21 @@ class PassiveTest(absltest.TestCase): mx = mjx.put_model(m) dx = jax.jit(mjx.passive)(mx, mjx.put_data(m, d)) _assert_attr_eq(d, dx, 'qfrc_passive') + _assert_attr_eq(d, dx, 'qfrc_gravcomp') m.opt.viscosity = 0.02 mujoco.mj_forward(m, d) mx = mjx.put_model(m) dx = jax.jit(mjx.passive)(mx, mjx.put_data(m, d)) _assert_attr_eq(d, dx, 'qfrc_passive') + _assert_attr_eq(d, dx, 'qfrc_gravcomp') m.opt.wind = np.array([0.03, 0.04, 0.05]) mujoco.mj_forward(m, d) mx = mjx.put_model(m) dx = jax.jit(mjx.passive)(mx, mjx.put_data(m, d)) _assert_attr_eq(d, dx, 'qfrc_passive') + _assert_attr_eq(d, dx, 'qfrc_gravcomp') # test disable passive mx = mx.tree_replace({'opt.disableflags': mjx.DisableBit.PASSIVE}) diff --git a/mjx/mujoco/mjx/_src/test_util.py b/mjx/mujoco/mjx/_src/test_util.py index a02d3922..07a45ba5 100644 --- a/mjx/mujoco/mjx/_src/test_util.py +++ b/mjx/mujoco/mjx/_src/test_util.py @@ -152,6 +152,7 @@ def _make_joint(joint_type: str, name: str) -> Dict[str, str]: joint_attr['damping'] = '{:.2f}'.format(np.random.uniform() * 20) joint_attr['stiffness'] = '{:.2f}'.format(np.random.uniform() * 20) + joint_attr['actuatorgravcomp'] = np.random.choice(['true', 'false']) return joint_attr @@ -334,7 +335,16 @@ def create_mjcf( z_pos = np.random.uniform(low=-1, high=1) * 0.01 # small jitter pos = f'{body_pos[0]:.3f} {body_pos[1]:.3f} {body_pos[2] + z_pos:.3f}' n_bodies = len(list(mjcf.iter('body'))) - child = ET.SubElement(body, 'body', {'pos': pos, 'name': f'body{n_bodies}'}) + gravcomp = np.random.uniform() * p(50) + child = ET.SubElement( + body, + 'body', + { + 'pos': pos, + 'name': f'body{n_bodies}', + 'gravcomp': f'{gravcomp:.3f}', + }, + ) ET.SubElement(child, 'site', {'name': f'site{n_bodies}'}) n_joints = len(list(mjcf.iter('joint'))) diff --git a/mjx/mujoco/mjx/_src/types.py b/mjx/mujoco/mjx/_src/types.py index 65f8e356..caccbea8 100644 --- a/mjx/mujoco/mjx/_src/types.py +++ b/mjx/mujoco/mjx/_src/types.py @@ -326,6 +326,7 @@ class Model(PyTreeNode): npair: number of predefined geom pairs nexclude: number of excluded geom pairs neq: number of equality constraints + ngravcomp: number of bodies with nonzero gravcomp nnumeric: number of numeric custom fields ntuple: number of tuple custom fields nsensor: number of sensors @@ -352,6 +353,7 @@ class Model(PyTreeNode): body_mass: mass (nbody,) body_subtreemass: mass of subtree starting at this body (nbody,) body_inertia: diagonal inertia in ipos/iquat frame (nbody, 3) + body_gravcomp: antigravity force, units of body weight (nbody,) body_invweight0: mean inv inert in qpos0 (trn, rot) (nbody, 2) jnt_type: type of joint (mjtJoint) (njnt,) jnt_qposadr: start addr in 'qpos' for joint's data (njnt,) @@ -359,6 +361,8 @@ class Model(PyTreeNode): jnt_bodyid: id of joint's body (njnt,) jnt_group: group for visibility (njnt,) jnt_limited: does joint have limits (njnt,) + jnt_actfrclimited: does joint have actuator force limits (njnt,) + jnt_actgravcomp: is gravcomp force applied via actuators (njnt,) jnt_solref: constraint solver reference: limit (njnt, mjNREF) jnt_solimp: constraint solver impedance: limit (njnt, mjNIMP) jnt_pos: local anchor position (njnt, 3) @@ -488,6 +492,7 @@ class Model(PyTreeNode): npair: int nexclude: int neq: int + ngravcomp: int nnumeric: int nuserdata: int ntuple: int @@ -514,6 +519,7 @@ class Model(PyTreeNode): body_mass: jax.Array body_subtreemass: jax.Array body_inertia: jax.Array + body_gravcomp: jax.Array body_invweight0: jax.Array jnt_type: np.ndarray jnt_qposadr: np.ndarray @@ -521,6 +527,7 @@ class Model(PyTreeNode): jnt_bodyid: np.ndarray jnt_limited: np.ndarray jnt_actfrclimited: np.ndarray + jnt_actgravcomp: np.ndarray jnt_solref: jax.Array jnt_solimp: jax.Array jnt_pos: jax.Array @@ -671,7 +678,7 @@ class Contact(PyTreeNode): class Data(PyTreeNode): - r"""Dynamic state that updates each step. + r"""\Dynamic state that updates each step. Attributes: ne: number of equality constraints @@ -725,6 +732,7 @@ class Data(PyTreeNode): cvel: com-based velocity [3D rot; 3D tran] (nbody, 6) cdof_dot: time-derivative of cdof (nv, 6) qfrc_bias: C(qpos,qvel) (nv,) + qfrc_gravcomp: passive gravity compensation force (nv,) qfrc_passive: passive force (nv,) efc_aref: reference pseudo-acceleration (nefc,) qfrc_actuator: actuator force (nv,) @@ -795,6 +803,7 @@ class Data(PyTreeNode): cdof_dot: jax.Array qfrc_bias: jax.Array qfrc_passive: jax.Array + qfrc_gravcomp: jax.Array efc_aref: jax.Array # position, velcoity, control & acceleration dependent: qfrc_actuator: jax.Array diff --git a/mjx/mujoco/mjx/integration_test/smooth_test.py b/mjx/mujoco/mjx/integration_test/smooth_test.py index c4d3bfda..5359245e 100644 --- a/mjx/mujoco/mjx/integration_test/smooth_test.py +++ b/mjx/mujoco/mjx/integration_test/smooth_test.py @@ -68,10 +68,10 @@ class TransmissionIntegrationTest(parameterized.TestCase): mujoco.mj_transmission(m, d) dx = transmission_jit_fn(mx, dx) - _assert_attr_eq(d, dx, 'actuator_length', seed, f'transmission{seed}') - _assert_attr_eq( - d, dx, 'actuator_moment', seed, f'transmission{seed}', atol=1e-4 - ) + for field in ['actuator_length', 'actuator_moment']: + _assert_attr_eq( + d, dx, field, seed, f'transmission{seed}', atol=1e-4 + ) if __name__ == '__main__': diff --git a/mjx/mujoco/mjx/test_data/pendula.xml b/mjx/mujoco/mjx/test_data/pendula.xml index f2a3fc8e..aff7741a 100644 --- a/mjx/mujoco/mjx/test_data/pendula.xml +++ b/mjx/mujoco/mjx/test_data/pendula.xml @@ -110,6 +110,20 @@ + + + + + + + + + + + + + + @@ -118,5 +132,8 @@ + + +