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):