Add site transmission to MJX.

PiperOrigin-RevId: 592373042
Change-Id: I54d941f4ee6f74404fe2aea7253f5b426d097e23
This commit is contained in:
Baruch Tabanpour
2023-12-19 16:18:01 -08:00
committed by Copybara-Service
parent 80f50c943c
commit 77b4132c47
10 changed files with 171 additions and 52 deletions
+1
View File
@@ -9,6 +9,7 @@ MJX
^^^
1. Add :ref:`dyntype<actuator-general-dyntype>` ``filterexact``.
2. Add :at:`site` transmission.
Version 3.1.1 (December 18, 2023)
+2 -2
View File
@@ -181,7 +181,7 @@ The following features are **fully supported** in MJX:
* - :ref:`Joint <mjtJoint>`
- ``FREE``, ``BALL``, ``SLIDE``, ``HINGE``
* - :ref:`Transmission <mjtTrn>`
- ``TRN_JOINT``
- ``TRN_JOINT``, ``TRN_SITE``
* - :ref:`Actuator Dynamics <mjtDyn>`
- ``NONE``, ``INTEGRATOR``, ``FILTER``, ``FILTEREXACT``
* - :ref:`Actuator Gain <mjtGain>`
@@ -218,7 +218,7 @@ The following features are **in development** and coming soon:
* - Dynamics
- :ref:`Inverse <mj_inverse>`
* - :ref:`Transmission <mjtTrn>`
- ``TRN_SITE``, ``TRN_TENDON``
- ``TRN_TENDON``
* - :ref:`Actuator Dynamics <mjtDyn>`
- ``MUSCLE``
* - :ref:`Actuator Gain <mjtGain>`
-6
View File
@@ -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
+4
View File
@@ -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)
+33 -8
View File
@@ -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("""
<mujoco>
@@ -174,7 +175,7 @@ class IoTest(parameterized.TestCase):
</actuator>
</mujoco>"""))
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("""
<mujoco>
@@ -191,7 +192,7 @@ class IoTest(parameterized.TestCase):
</tendon>
</mujoco>"""))
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("""
<mujoco>
@@ -207,7 +208,7 @@ class IoTest(parameterized.TestCase):
</worldbody>
</mujoco>"""))
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("""
<mujoco>
@@ -223,6 +224,29 @@ class IoTest(parameterized.TestCase):
</worldbody>
</mujoco>"""))
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("""
<mujoco>
<compiler autolimits="true"/>
<worldbody>
<body name="box">
<site name="site1"/>
<site name="site2" pos="0.2 0.1 0.05"/>
<joint name="slide" type="slide" axis="1 0 0" />
<geom type="box" size=".05 .05 .05" mass="1"/>
</body>
</worldbody>
<actuator>
<position site="site2" refsite="site1"/>
</actuator>
</mujoco>"""))
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()
+31 -9
View File
@@ -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
+37 -18
View File
@@ -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,))
+35
View File
@@ -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("""
<mujoco>
<compiler autolimits="true"/>
<worldbody>
<body>
<joint type="free"/>
<geom type="box" size=".05 .05 .05" mass="1"/>
<site name="site1"/>
<site name="site2" pos="0.1 0.2 0.3"/>
</body>
<body pos="1 0 0">
<joint name="slide" type="hinge"/>
<geom type="box" size=".05 .05 .05" mass="1"/>
</body>
</worldbody>
<actuator>
<position site="site1" gear="1 2 3 0 0 0"/>
<position site="site1" gear="0 0 0 1 2 3"/>
<position site="site2" gear="0 3 0 0 0 1"/>
<position joint="slide"/>
</actuator>
</mujoco>
""")
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()
+25 -8
View File
@@ -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
+3 -1
View File
@@ -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):