Add site transmission to MJX.
PiperOrigin-RevId: 592373042 Change-Id: I54d941f4ee6f74404fe2aea7253f5b426d097e23
This commit is contained in:
committed by
Copybara-Service
parent
80f50c943c
commit
77b4132c47
@@ -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
@@ -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>`
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,))
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user