From 8ce2c92021ed202db6bbb71bf4e30e4f81150b4c Mon Sep 17 00:00:00 2001 From: Baruch Tabanpour Date: Tue, 2 Jan 2024 16:44:59 -0800 Subject: [PATCH] Add refsite transmission. PiperOrigin-RevId: 595239860 Change-Id: Id1ca8fe2ffab4b55013e8987e92eb99847278dc4 --- mjx/mujoco/mjx/_src/dataclasses.py | 4 +- mjx/mujoco/mjx/_src/io.py | 5 -- mjx/mujoco/mjx/_src/io_test.py | 35 --------- mjx/mujoco/mjx/_src/scan.py | 5 +- mjx/mujoco/mjx/_src/smooth.py | 73 +++++++++++++++++-- mjx/mujoco/mjx/_src/smooth_test.py | 17 +++-- mjx/mujoco/mjx/_src/test_util.py | 18 ++++- mjx/mujoco/mjx/_src/types.py | 2 - .../mjx/integration_test/smooth_test.py | 6 +- 9 files changed, 103 insertions(+), 62 deletions(-) diff --git a/mjx/mujoco/mjx/_src/dataclasses.py b/mjx/mujoco/mjx/_src/dataclasses.py index b939ab9e..5d95eb69 100644 --- a/mjx/mujoco/mjx/_src/dataclasses.py +++ b/mjx/mujoco/mjx/_src/dataclasses.py @@ -62,7 +62,7 @@ def dataclass(clz: _T) -> _T: def to_meta(field, obj): val = getattr(obj, field.name) - return to_tup(val) if isinstance(val, np.ndarray) else val + return (to_tup(val), val.dtype) if isinstance(val, np.ndarray) else val def to_data(field, obj): return (jax.tree_util.GetAttrKey(field.name), getattr(obj, field.name)) @@ -75,7 +75,7 @@ def dataclass(clz: _T) -> _T: def from_meta(field, meta): if field.type is np.ndarray: - return (field.name, np.array(meta)) + return (field.name, np.array(meta[0], dtype=meta[1])) else: return (field.name, meta) diff --git a/mjx/mujoco/mjx/_src/io.py b/mjx/mujoco/mjx/_src/io.py index ab255cfa..788c3899 100644 --- a/mjx/mujoco/mjx/_src/io.py +++ b/mjx/mujoco/mjx/_src/io.py @@ -103,10 +103,6 @@ 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) @@ -196,7 +192,6 @@ def make_data(m: Union[types.Model, mujoco.MjModel]) -> types.Data: qfrc_bias=zero_nv, qfrc_passive=zero_nv, efc_aref=zero_nefc, - actuator_force=zero_nu, qfrc_actuator=zero_nv, qfrc_smooth=zero_nv, qacc_smooth=zero_nv, diff --git a/mjx/mujoco/mjx/_src/io_test.py b/mjx/mujoco/mjx/_src/io_test.py index 08dec57a..8b25f359 100644 --- a/mjx/mujoco/mjx/_src/io_test.py +++ b/mjx/mujoco/mjx/_src/io_test.py @@ -160,21 +160,6 @@ class ModelIOTest(parameterized.TestCase): ) ) - def test_site_actuator_not_implemented(self): - with self.assertRaises(NotImplementedError): - mjx.put_model(mujoco.MjModel.from_xml_string(""" - - - - - - - - - - - """)) - def test_tendon_not_implemented(self): with self.assertRaises(NotImplementedError): mjx.put_model(mujoco.MjModel.from_xml_string(""" @@ -224,25 +209,6 @@ class ModelIOTest(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.""" @@ -305,7 +271,6 @@ class DataIOTest(parameterized.TestCase): self.assertEqual(d.qfrc_bias.shape, (nv,)) self.assertEqual(d.qfrc_passive.shape, (nv,)) self.assertEqual(d.efc_aref.shape, (nefc,)) - self.assertEqual(d.actuator_force.shape, (1,)) self.assertEqual(d.qfrc_actuator.shape, (nv,)) self.assertEqual(d.qfrc_smooth.shape, (nv,)) self.assertEqual(d.qacc_smooth.shape, (nv,)) diff --git a/mjx/mujoco/mjx/_src/scan.py b/mjx/mujoco/mjx/_src/scan.py index 9035c98e..89f360cb 100644 --- a/mjx/mujoco/mjx/_src/scan.py +++ b/mjx/mujoco/mjx/_src/scan.py @@ -220,6 +220,7 @@ def flat( m.actuator_dyntype[ids_u], m.actuator_trntype[ids_u], m.jnt_type[ids_j], + m.actuator_trnid[ids_u, 1] == -1, # key by refsite being present ) def type_ids_j(m, i): @@ -240,9 +241,9 @@ def flat( else -1 ), 's': ( - m.actuator_trnid[i, 0] + m.actuator_trnid[i] if m.actuator_trntype[i] == TrnType.SITE - else -1 + else np.array([-1, -1]) ), } v, q = np.array([-1]), np.array([-1]) diff --git a/mjx/mujoco/mjx/_src/smooth.py b/mjx/mujoco/mjx/_src/smooth.py index 81c53213..c1531192 100644 --- a/mjx/mujoco/mjx/_src/smooth.py +++ b/mjx/mujoco/mjx/_src/smooth.py @@ -27,6 +27,7 @@ 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 +import numpy as np def kinematics(m: Model, d: Data) -> Data: @@ -432,13 +433,54 @@ def rne(m: Model, d: Data) -> Data: return d +def _site_dof_mask(m: Model) -> np.ndarray: + """Creates a dof mask for site transmissions.""" + mask = np.ones((m.nu, m.nv)) + for i in np.nonzero(m.actuator_trnid[:, 1] != -1)[0]: + id_, refid = m.actuator_trnid[i] + # intialize last dof address for each body + b0 = m.body_weldid[m.site_bodyid[id_]] + b1 = m.body_weldid[m.site_bodyid[refid]] + dofadr0 = m.body_dofadr[b0] + m.body_dofnum[b0] - 1 + dofadr1 = m.body_dofadr[b1] + m.body_dofnum[b1] - 1 + + # find common ancestral dof, if any + while dofadr0 != dofadr1: + if dofadr0 < dofadr1: + dofadr1 = m.dof_parentid[dofadr1] + else: + dofadr0 = m.dof_parentid[dofadr0] + if dofadr0 == -1 or dofadr1 == -1: + break + + # if common ancestral dof was found, clear the columns of its parental chain + da = dofadr0 if dofadr0 == dofadr1 else -1 + while da >= 0: + mask[i, da] = 0.0 + da = m.dof_parentid[da] + + return mask + + 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(trntype, trnid, gear, jnt_typ, m_j, qpos, site_xpos, site_xmat): + def fn( + trntype, + trnid, + gear, + jnt_typ, + m_j, + qpos, + has_refsite, + site_dof_mask, + site_xpos, + site_xmat, + site_quat, + ): if trntype == TrnType.JOINT: if jnt_typ == JointType.FREE: length = jp.zeros(1) @@ -459,21 +501,34 @@ def transmission(m: Model, d: Data) -> Data: moment = jp.zeros((m.nv,)).at[m_j].set(moment) elif trntype == TrnType.SITE: length = jp.zeros(1) - 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:])) + id_, refid = jp.array(m.site_bodyid)[trnid] + jacp, jacr = support.jac(m, d, site_xpos[0], id_) + frame_xmat = site_xmat[0] + if has_refsite: + vecp = site_xmat[1].T @ (site_xpos[0] - site_xpos[1]) + vecr = math.quat_sub(site_quat[0], site_quat[1]) + length += jp.dot(jp.concatenate([vecp, vecr]), gear) + jacrefp, jacrefr = support.jac(m, d, site_xpos[1], refid) + jacp, jacr = jacp - jacrefp, jacr - jacrefr + frame_xmat = site_xmat[1] + + jac = jp.concatenate((jacp, jacr), axis=1) * site_dof_mask[:, None] + wrench = jp.concatenate((frame_xmat @ gear[:3], frame_xmat @ gear[3:])) moment = jac @ wrench else: raise RuntimeError(f'unrecognized trntype: {TrnType(trntype)}') return length, moment + # pre-compute values for site transmissions + has_refsite = m.actuator_trnid[:, 1] != -1 + site_dof_mask = _site_dof_mask(m) + site_quat = jax.vmap(math.quat_mul)(m.site_quat, d.xquat[m.site_bodyid]) + length, moment = scan.flat( m, fn, - 'uuujjqss', + 'uuujjquusss', 'uu', m.actuator_trntype, jp.array(m.actuator_trnid), @@ -481,11 +536,15 @@ def transmission(m: Model, d: Data) -> Data: m.jnt_type, jp.array(m.jnt_dofadr), d.qpos, + has_refsite, + jp.array(site_dof_mask), d.site_xpos, d.site_xmat, + site_quat, group_by='u', ) length = length.reshape((m.nu,)) moment = moment.reshape((m.nu, m.nv)) + d = d.replace(actuator_length=length, actuator_moment=moment) return d diff --git a/mjx/mujoco/mjx/_src/smooth_test.py b/mjx/mujoco/mjx/_src/smooth_test.py index e83779ad..b7d62eb9 100644 --- a/mjx/mujoco/mjx/_src/smooth_test.py +++ b/mjx/mujoco/mjx/_src/smooth_test.py @@ -141,7 +141,11 @@ class SmoothTest(absltest.TestCase): - + + + + + @@ -149,10 +153,11 @@ class SmoothTest(absltest.TestCase): - - - - + + + + + """) @@ -162,7 +167,7 @@ class SmoothTest(absltest.TestCase): dx = mjx.put_data(m, d) mujoco.mj_transmission(m, d) - dx = mjx.transmission(mx, dx) + dx = jax.jit(mjx.transmission)(mx, dx) _assert_attr_eq(d, dx, 'actuator_length') _assert_attr_eq(d, dx, 'actuator_moment') diff --git a/mjx/mujoco/mjx/_src/test_util.py b/mjx/mujoco/mjx/_src/test_util.py index 66bab2f7..ab27c2b1 100644 --- a/mjx/mujoco/mjx/_src/test_util.py +++ b/mjx/mujoco/mjx/_src/test_util.py @@ -36,7 +36,7 @@ _JOINT_AXES = ['1 0 0', '0 1 0', '0 0 1'] _FRICTIONS = ['1.2 0.003 0.0002', '0.2 0.0001 0.0005'] _KP_POS = ['1', '2'] _KP_INTVEL = ['10000', '2000'] -_KV_VEL = ['123', '1'] +_KV_VEL = ['12', '1', '0', '0.1'] _PAIR_FRICTIONS = ['1.2 0.9 0.003 0.0002 0.0001'] _SOLREFS = ['0.04 1.01', '0.05 1.02', '0.03 1.1', '0.015 1.0'] _SOLIMPS = [ @@ -124,7 +124,10 @@ def _make_geom( def _make_actuator( - actuator_type: str, joint: str | None = None, site: str | None = None + actuator_type: str, + joint: str | None = None, + site: str | None = None, + refsite: str | None = None, ) -> Dict[str, str]: """Returns attributes for an actuator.""" if joint: @@ -134,11 +137,15 @@ def _make_actuator( else: raise ValueError('must provide a joint or site name') + if refsite: + attr['refsite'] = refsite + attr['gear'] = np.random.choice(_GEARS) # set actuator type if actuator_type == 'position': attr['kp'] = np.random.choice(_KP_POS) + attr['kv'] = np.random.choice(_KV_VEL) elif actuator_type == 'general': attr['biastype'] = 'affine' attr['gainprm'] = '35 0 0' @@ -312,6 +319,13 @@ def create_mjcf( attr = _make_actuator(actuator_type, site=f'site{i}') actuators.append((actuator_type, attr)) + # site transmission with refsite + for i in range(np.random.randint(0, n_bodies)): + j = np.random.randint(0, n_bodies) + actuator_type = np.random.choice(_ACTUATOR_TYPES) + attr = _make_actuator(actuator_type, site=f'site{i}', refsite=f'site{j}') + actuators.append((actuator_type, attr)) + np.random.shuffle(actuators) for typ, attr in actuators: ET.SubElement(actuator, typ, attr) diff --git a/mjx/mujoco/mjx/_src/types.py b/mjx/mujoco/mjx/_src/types.py index ece81623..595d7f90 100644 --- a/mjx/mujoco/mjx/_src/types.py +++ b/mjx/mujoco/mjx/_src/types.py @@ -590,7 +590,6 @@ class Data(PyTreeNode): qfrc_bias: C(qpos,qvel) (nv,) qfrc_passive: passive force (nv,) efc_aref: reference pseudo-acceleration (nefc,) - actuator_force: actuator force in actuation space (nu,) qfrc_actuator: actuator force (nv,) qfrc_smooth: net unconstrained force (nv,) qacc_smooth: unconstrained acceleration (nv,) @@ -650,7 +649,6 @@ class Data(PyTreeNode): qfrc_passive: jax.Array efc_aref: jax.Array # position, velcoity, control & acceleration dependent: - actuator_force: jax.Array qfrc_actuator: jax.Array qfrc_smooth: jax.Array qacc_smooth: jax.Array diff --git a/mjx/mujoco/mjx/integration_test/smooth_test.py b/mjx/mujoco/mjx/integration_test/smooth_test.py index d032be79..c4d3bfda 100644 --- a/mjx/mujoco/mjx/integration_test/smooth_test.py +++ b/mjx/mujoco/mjx/integration_test/smooth_test.py @@ -57,7 +57,9 @@ class TransmissionIntegrationTest(parameterized.TestCase): d = mujoco.MjData(m) d.ctrl = np.random.normal(scale=10, size=m.nu) d.act = np.random.normal(scale=10, size=m.na) + d.qpos = np.random.normal(m.nq) d.qvel = np.random.random(m.nv) + mujoco.mj_forward(m, d) # put on device mx = mjx.put_model(m) @@ -67,7 +69,9 @@ class TransmissionIntegrationTest(parameterized.TestCase): 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}') + _assert_attr_eq( + d, dx, 'actuator_moment', seed, f'transmission{seed}', atol=1e-4 + ) if __name__ == '__main__':