diff --git a/doc/changelog.rst b/doc/changelog.rst index ad621fb4..52592511 100644 --- a/doc/changelog.rst +++ b/doc/changelog.rst @@ -9,6 +9,7 @@ MJX ^^^ 1. Add :ref:`dyntype` ``filterexact``. +2. Add :at:`site` transmission. Version 3.1.1 (December 18, 2023) diff --git a/doc/mjx.rst b/doc/mjx.rst index cc052182..c9263fa9 100644 --- a/doc/mjx.rst +++ b/doc/mjx.rst @@ -181,7 +181,7 @@ The following features are **fully supported** in MJX: * - :ref:`Joint ` - ``FREE``, ``BALL``, ``SLIDE``, ``HINGE`` * - :ref:`Transmission ` - - ``TRN_JOINT`` + - ``TRN_JOINT``, ``TRN_SITE`` * - :ref:`Actuator Dynamics ` - ``NONE``, ``INTEGRATOR``, ``FILTER``, ``FILTEREXACT`` * - :ref:`Actuator Gain ` @@ -218,7 +218,7 @@ The following features are **in development** and coming soon: * - Dynamics - :ref:`Inverse ` * - :ref:`Transmission ` - - ``TRN_SITE``, ``TRN_TENDON`` + - ``TRN_TENDON`` * - :ref:`Actuator Dynamics ` - ``MUSCLE`` * - :ref:`Actuator Gain ` diff --git a/mjx/mujoco/mjx/_src/device_test.py b/mjx/mujoco/mjx/_src/device_test.py index 9f518d91..b4b9a0cd 100644 --- a/mjx/mujoco/mjx/_src/device_test.py +++ b/mjx/mujoco/mjx/_src/device_test.py @@ -129,12 +129,6 @@ class ValidateInputTest(absltest.TestCase): with self.assertRaises(NotImplementedError): mjx.device_put(m) - def test_trn(self): - m = test_util.load_test_file('pendula.xml') - m.actuator_trntype[0] = mujoco.mjtTrn.mjTRN_SITE - with self.assertRaises(NotImplementedError): - mjx.device_put(m) - def test_dyn(self): m = test_util.load_test_file('pendula.xml') m.actuator_dyntype[0] = mujoco.mjtDyn.mjDYN_MUSCLE diff --git a/mjx/mujoco/mjx/_src/io.py b/mjx/mujoco/mjx/_src/io.py index 9396d7f3..ab255cfa 100644 --- a/mjx/mujoco/mjx/_src/io.py +++ b/mjx/mujoco/mjx/_src/io.py @@ -103,6 +103,10 @@ def put_model(m: mujoco.MjModel, device=None) -> types.Model: f'{[mj_type(m) for m in missing]} not supported' ) + # TODO: implement reference sites. + if any(m.actuator_trnid[:, 1] != -1): + raise NotImplementedError('refsite is not supported') + opt = _put_option(m.opt, device=device) stat = _put_statistic(m.stat, device=device) diff --git a/mjx/mujoco/mjx/_src/io_test.py b/mjx/mujoco/mjx/_src/io_test.py index a16a8780..08dec57a 100644 --- a/mjx/mujoco/mjx/_src/io_test.py +++ b/mjx/mujoco/mjx/_src/io_test.py @@ -82,7 +82,8 @@ _MULTIPLE_CONSTRAINTS = """ """ -class IoTest(parameterized.TestCase): +class ModelIOTest(parameterized.TestCase): + """IO tests for mjx.Model.""" def test_put_model(self): m = mujoco.MjModel.from_xml_string(_MULTIPLE_CONVEX_OBJECTS) @@ -133,7 +134,7 @@ class IoTest(parameterized.TestCase): ) self.assertTrue(m.opt.has_fluid_params) - def test_put_model_implicit_not_implemented(self): + def test_implicit_not_implemented(self): """Test that MJX guards against models with unimplemented features.""" with self.assertRaises(NotImplementedError): @@ -143,7 +144,7 @@ class IoTest(parameterized.TestCase): ) ) - def test_put_model_cone_not_implemented(self): + def test_cone_not_implemented(self): with self.assertRaises(NotImplementedError): mjx.put_model( mujoco.MjModel.from_xml_string( @@ -151,7 +152,7 @@ class IoTest(parameterized.TestCase): ) ) - def test_put_model_pgs_not_implemented(self): + def test_pgs_not_implemented(self): with self.assertRaises(NotImplementedError): mjx.put_model( mujoco.MjModel.from_xml_string( @@ -159,7 +160,7 @@ class IoTest(parameterized.TestCase): ) ) - def test_put_model_site_actuator_not_implemented(self): + def test_site_actuator_not_implemented(self): with self.assertRaises(NotImplementedError): mjx.put_model(mujoco.MjModel.from_xml_string(""" @@ -174,7 +175,7 @@ class IoTest(parameterized.TestCase): """)) - def test_put_model_tendon_not_implemented(self): + def test_tendon_not_implemented(self): with self.assertRaises(NotImplementedError): mjx.put_model(mujoco.MjModel.from_xml_string(""" @@ -191,7 +192,7 @@ class IoTest(parameterized.TestCase): """)) - def test_put_model_condim_not_implemented(self): + def test_condim_not_implemented(self): with self.assertRaises(NotImplementedError): mjx.put_model(mujoco.MjModel.from_xml_string(""" @@ -207,7 +208,7 @@ class IoTest(parameterized.TestCase): """)) - def test_put_model_cylinder_not_implemented(self): + def test_cylinder_not_implemented(self): with self.assertRaises(NotImplementedError): mjx.put_model(mujoco.MjModel.from_xml_string(""" @@ -223,6 +224,29 @@ class IoTest(parameterized.TestCase): """)) + def test_refsite_not_implemented(self): + """Tests that site transmissions with refsites are not implemented.""" + with self.assertRaises(NotImplementedError): + mjx.put_model(mujoco.MjModel.from_xml_string(""" + + + + + + + + + + + + + + """)) + + +class DataIOTest(parameterized.TestCase): + """IO tests for mjx.Data.""" + def test_make_data(self): """Test that make_data returns the correct shapes.""" @@ -407,5 +431,6 @@ class IoTest(parameterized.TestCase): self.assertEqual(ds[0].ncon, 1) self.assertEqual(ds[1].ncon, 0) + if __name__ == '__main__': absltest.main() diff --git a/mjx/mujoco/mjx/_src/scan.py b/mjx/mujoco/mjx/_src/scan.py index ae0eeb30..9035c98e 100644 --- a/mjx/mujoco/mjx/_src/scan.py +++ b/mjx/mujoco/mjx/_src/scan.py @@ -49,7 +49,9 @@ def _take(obj: Y, idx: np.ndarray) -> Y: def take(x): # TODO(erikfrey): if this helps perf, add support for striding too - if ( + if not x.shape[0]: + return x + elif ( len(idx.shape) == 1 and idx.size > 0 and (idx == np.arange(idx[0], idx[0] + idx.size)).all() @@ -113,8 +115,12 @@ def _nvmap(f: Callable[..., Y], *args) -> Y: if isinstance(arg, np.ndarray) and not np.all(arg == arg[0]): raise RuntimeError(f'numpy arg elements do not match: {arg}') + # split out numpy and jax args np_args = [a[0] if isinstance(a, np.ndarray) else None for a in args] args = [a if n is None else None for n, a in zip(np_args, args)] + + # remove empty args that we should not vmap over + args = jax.tree_map(lambda a: a if a.shape[0] else None, args) in_axes = [None if a is None else 0 for a in args] def outer_f(*args, np_args=np_args): @@ -126,7 +132,15 @@ def _nvmap(f: Callable[..., Y], *args) -> Y: def _check_input(m: Model, args: Any, in_types: str) -> None: """Checks that scan input has the right shape.""" - size = {'b': m.nbody, 'j': m.njnt, 'q': m.nq, 'v': m.nv, 'u': m.nu, 'a': m.na} + size = { + 'b': m.nbody, + 'j': m.njnt, + 'q': m.nq, + 'v': m.nv, + 'u': m.nu, + 'a': m.na, + 's': m.nsite, + } for idx, (arg, typ) in enumerate(zip(args, in_types)): if len(arg) != size[typ]: raise IndexError( @@ -162,7 +176,7 @@ def flat( ) -> Y: r"""Scan a function across bodies or actuators. - Scan group data according to type and batch shape then calls vmap(f) on it. + Scan group data according to type and batch shape then calls vmap(f) on it.\ Args: m: an mjx model @@ -223,14 +237,22 @@ def flat( 'j': ( m.actuator_trnid[i, 0] if m.actuator_trntype[i] == TrnType.JOINT - else np.array(-1) + else -1 + ), + 's': ( + m.actuator_trnid[i, 0] + if m.actuator_trntype[i] == TrnType.SITE + else -1 ), } - # v/q associated with joint transmissions - typ_ids.update({ - 'v': np.nonzero(m.dof_jntid == typ_ids['j'])[0], - 'q': np.nonzero(_q_jointid(m) == typ_ids['j'])[0], - }) + v, q = np.array([-1]), np.array([-1]) + if m.actuator_trntype[i] == TrnType.JOINT: + # v/q are associated with the joint transmissions only + v = np.nonzero(m.dof_jntid == typ_ids['j'])[0] + q = np.nonzero(_q_jointid(m) == typ_ids['j'])[0] + + typ_ids.update({'v': v, 'q': q}) + return typ_ids # build up a grouping of type take-ids in body/actuator order diff --git a/mjx/mujoco/mjx/_src/smooth.py b/mjx/mujoco/mjx/_src/smooth.py index ba470bed..81c53213 100644 --- a/mjx/mujoco/mjx/_src/smooth.py +++ b/mjx/mujoco/mjx/_src/smooth.py @@ -19,11 +19,13 @@ from jax import numpy as jp import mujoco from mujoco.mjx._src import math from mujoco.mjx._src import scan +from mujoco.mjx._src import support # pylint: disable=g-importing-member from mujoco.mjx._src.types import Data from mujoco.mjx._src.types import DisableBit from mujoco.mjx._src.types import JointType from mujoco.mjx._src.types import Model +from mujoco.mjx._src.types import TrnType # pylint: enable=g-importing-member @@ -432,38 +434,55 @@ def rne(m: Model, d: Data) -> Data: def transmission(m: Model, d: Data) -> Data: """Computes actuator/transmission lengths and moments.""" + # TODO: consider combining transmission calculation into fwd_actuation. if not m.nu: return d - def fn(gear, jnt_typ, m_j, qpos): - # handles joint transmissions only - if jnt_typ == JointType.FREE: + def fn(trntype, trnid, gear, jnt_typ, m_j, qpos, site_xpos, site_xmat): + if trntype == TrnType.JOINT: + if jnt_typ == JointType.FREE: + length = jp.zeros(1) + moment = gear + m_j = m_j + jp.arange(6) + elif jnt_typ == JointType.BALL: + axis, angle = math.quat_to_axis_angle(qpos) + length = jp.dot(axis * angle, gear[:3])[None] + moment = gear[:3] + m_j = m_j + jp.arange(3) + elif jnt_typ in (JointType.SLIDE, JointType.HINGE): + length = qpos * gear[0] + moment = gear[:1] + m_j = m_j[None] + else: + raise RuntimeError(f'unrecognized joint type: {JointType(jnt_typ)}') + + moment = jp.zeros((m.nv,)).at[m_j].set(moment) + elif trntype == TrnType.SITE: length = jp.zeros(1) - moment = gear - m_j = m_j + jp.arange(6) - elif jnt_typ == JointType.BALL: - axis, angle = math.quat_to_axis_angle(qpos) - length = jp.dot(axis * angle, gear[:3])[None] - moment = gear[:3] - m_j = m_j + jp.arange(3) - elif jnt_typ in (JointType.SLIDE, JointType.HINGE): - length = qpos * gear[0] - moment = gear[:1] - m_j = m_j[None] + jacp, jacr = support.jac( + m, d, site_xpos, jp.array(m.site_bodyid)[trnid[0]] + ) + jac = jp.concatenate((jacp, jacr), axis=1) + wrench = jp.concatenate((site_xmat @ gear[:3], site_xmat @ gear[3:])) + moment = jac @ wrench else: - raise RuntimeError(f'unrecognized joint type: {jnt_typ}') - moment = jp.zeros((m.nv,)).at[m_j].set(moment) + raise RuntimeError(f'unrecognized trntype: {TrnType(trntype)}') + return length, moment length, moment = scan.flat( m, fn, - 'ujjq', - 'uuuu', + 'uuujjqss', + 'uu', + m.actuator_trntype, + jp.array(m.actuator_trnid), m.actuator_gear, m.jnt_type, jp.array(m.jnt_dofadr), d.qpos, + d.site_xpos, + d.site_xmat, group_by='u', ) length = length.reshape((m.nu,)) diff --git a/mjx/mujoco/mjx/_src/smooth_test.py b/mjx/mujoco/mjx/_src/smooth_test.py index 203cb654..e83779ad 100644 --- a/mjx/mujoco/mjx/_src/smooth_test.py +++ b/mjx/mujoco/mjx/_src/smooth_test.py @@ -132,5 +132,40 @@ class SmoothTest(absltest.TestCase): dx = jax.jit(mjx.rne)(mx, dx) np.testing.assert_allclose(dx.qfrc_bias, 0) + def test_site_transmission(self): + m = mujoco.MjModel.from_xml_string(""" + + + + + + + + + + + + + + + + + + + + + + """) + d = mujoco.MjData(m) + mujoco.mj_forward(m, d) + mx = mjx.put_model(m) + dx = mjx.put_data(m, d) + + mujoco.mj_transmission(m, d) + dx = mjx.transmission(mx, dx) + _assert_attr_eq(d, dx, 'actuator_length') + _assert_attr_eq(d, dx, 'actuator_moment') + + if __name__ == '__main__': absltest.main() diff --git a/mjx/mujoco/mjx/_src/test_util.py b/mjx/mujoco/mjx/_src/test_util.py index 85cdb5dc..66bab2f7 100644 --- a/mjx/mujoco/mjx/_src/test_util.py +++ b/mjx/mujoco/mjx/_src/test_util.py @@ -47,7 +47,7 @@ _SOLIMPS = [ _DIMS = ['3'] _MARGINS = ['0.0', '0.01', '0.02'] _GAPS = ['0.0', '0.005'] -_GEARS = ['20', '50', '100'] +_GEARS = ['2.1 0.0 3.3 0 2.3 0', '5.0 3.1 0 2.3 0.0 1.1'] def p(pct: int) -> bool: @@ -123,14 +123,21 @@ def _make_geom( return attr -def _make_actuator(actuator_type: str, joint: str) -> Dict[str, str]: +def _make_actuator( + actuator_type: str, joint: str | None = None, site: str | None = None +) -> Dict[str, str]: """Returns attributes for an actuator.""" - attr = {'joint': joint} + if joint: + attr = {'joint': joint} + elif site: + attr = {'site': site} + else: + raise ValueError('must provide a joint or site name') + + attr['gear'] = np.random.choice(_GEARS) # set actuator type - if actuator_type == 'motor': - attr['gear'] = np.random.choice(_GEARS) - elif actuator_type == 'position': + if actuator_type == 'position': attr['kp'] = np.random.choice(_KP_POS) elif actuator_type == 'general': attr['biastype'] = 'affine' @@ -245,6 +252,7 @@ def create_mjcf( 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}'}) + ET.SubElement(child, 'site', {'name': f'site{n_bodies}'}) n_joints = len(list(mjcf.iter('joint'))) for nj in range(np.random.randint(1, max_stacked_joints + 1)): @@ -282,17 +290,28 @@ def create_mjcf( for _ in range(num_trees): make_tree(world, 0) + bodies = list(mjcf.iter('body')) + n_bodies = len(bodies) + # actuators if add_actuators: actuator = ET.SubElement(mjcf, 'actuator') n_joints = len(list(mjcf.iter('joint'))) nu = np.random.randint(1, n_joints + 1) actuators = [] + + # joint transmission for i in range(nu): actuator_type = np.random.choice(_ACTUATOR_TYPES) attr = _make_actuator(actuator_type, joint=f'joint{i}') actuators.append((actuator_type, attr)) + # site transmission + for i in range(np.random.randint(0, n_bodies)): + actuator_type = np.random.choice(_ACTUATOR_TYPES) + attr = _make_actuator(actuator_type, site=f'site{i}') + actuators.append((actuator_type, attr)) + np.random.shuffle(actuators) for typ, attr in actuators: ET.SubElement(actuator, typ, attr) @@ -320,9 +339,7 @@ def create_mjcf( ET.SubElement(contact, 'pair', attr) # exclude contacts - bodies = list(mjcf.iter('body')) body_names = [b.get('name') for b in bodies] - n_bodies = len(bodies) for _ in range(min(max_contact_excludes, (n_bodies * (n_bodies - 1) // 2))): if p(50): continue diff --git a/mjx/mujoco/mjx/_src/types.py b/mjx/mujoco/mjx/_src/types.py index 9ccbe29a..ece81623 100644 --- a/mjx/mujoco/mjx/_src/types.py +++ b/mjx/mujoco/mjx/_src/types.py @@ -153,9 +153,11 @@ class TrnType(enum.IntEnum): Attributes: JOINT: force on joint + SITE: force on site """ JOINT = mujoco.mjtTrn.mjTRN_JOINT - # unsupported: JOINTINPARENT, SLIDERCRANK, TENDON, SITE, BODY + SITE = mujoco.mjtTrn.mjTRN_SITE + # unsupported: JOINTINPARENT, SLIDERCRANK, TENDON, BODY class DynType(enum.IntEnum):