Add subtree_vel function to MJX. This function matches mj_subtreeVel.

PiperOrigin-RevId: 666078305
Change-Id: Ia0ecbf8279e914f49ee60b862541d697e8d4abb1
This commit is contained in:
Taylor Howell
2024-08-21 16:20:39 -07:00
committed by Copybara-Service
parent e76168141a
commit b42780a5f4
3 changed files with 113 additions and 0 deletions
+1
View File
@@ -43,6 +43,7 @@ from mujoco.mjx._src.smooth import crb
from mujoco.mjx._src.smooth import factor_m
from mujoco.mjx._src.smooth import kinematics
from mujoco.mjx._src.smooth import rne
from mujoco.mjx._src.smooth import subtree_vel
from mujoco.mjx._src.smooth import tendon
from mujoco.mjx._src.smooth import transmission
from mujoco.mjx._src.solver import solve
+93
View File
@@ -427,6 +427,99 @@ def com_vel(m: Model, d: Data) -> Data:
return d
def subtree_vel(m: Model, d: Data) -> Data:
"""Subtree linear velocity and angular momentum."""
# bodywise quantities
def _forward(cvel, xipos, ximat, subtree_com_root, mass, inertia):
ang, lin = jp.split(cvel, 2)
# update linear velocity
lin = lin - jp.cross(xipos - subtree_com_root, ang)
subtree_linvel = mass * lin
subtree_angmom = inertia * ximat @ ximat.T @ ang
body_vel = jp.concatenate([ang, lin])
return body_vel, subtree_linvel, subtree_angmom
body_vel, subtree_linvel, subtree_angmom = jax.vmap(_forward)(
d.cvel,
d.xipos,
d.ximat,
d.subtree_com[m.body_rootid],
m.body_mass,
m.body_inertia,
)
# sum body linear momentum recursively up the kinematic tree
subtree_linvel = scan.body_tree(
m,
lambda x, y: y if x is None else x + y,
'bb',
'b',
subtree_linvel,
reverse=True,
)
subtree_linvel /= jp.maximum(mujoco.mjMINVAL, m.body_subtreemass)[:, None]
def _subtree_angmom(
carry,
angmom,
com,
com_parent,
linvel,
linvel_parent,
subtreemass,
xipos,
vel,
mass,
mask,
):
def _momentum(x0, x1, v0, v1, m):
dx = x0 - x1
dv = v0 - v1
dp = dv * m
return jp.cross(dx, dp)
# momentum wrt current body
mom = mask * _momentum(xipos, com, vel[3:], linvel, mass)
# momentum wrt parent
mom_parent = mask * _momentum(
com, com_parent, linvel, linvel_parent, subtreemass
)
if carry is None:
return angmom + mom, mom_parent
else:
angmom_child, mom_parent_child = carry
return angmom + mom + angmom_child + mom_parent_child, mom_parent
subtree_angmom, _ = scan.body_tree(
m,
_subtree_angmom,
'bbbbbbbbbb',
'bb',
subtree_angmom,
d.subtree_com,
d.subtree_com[m.body_parentid],
subtree_linvel,
subtree_linvel[m.body_parentid],
m.body_subtreemass,
d.xipos,
body_vel,
m.body_mass,
jp.ones(m.nbody).at[0].set(0),
reverse=True,
)
return d.replace(subtree_linvel=subtree_linvel, subtree_angmom=subtree_angmom)
def rne(m: Model, d: Data) -> Data:
"""Computes inverse dynamics using the recursive Newton-Euler algorithm."""
# forward scan over tree: accumulate link center of mass acceleration
+19
View File
@@ -172,6 +172,25 @@ class SmoothTest(absltest.TestCase):
_assert_attr_eq(d, dx, 'actuator_length')
_assert_attr_eq(d, dx, 'actuator_moment')
def test_subtree_vel(self):
"""Tests MJX subtree_vel function matches MuJoCo mj_subtreeVel."""
m = test_util.load_test_file('humanoid/humanoid.xml')
d = mujoco.MjData(m)
# give the system a little kick to ensure we have non-identity rotations
d.qvel = np.random.random(m.nv)
mujoco.mj_step(m, d, 10) # let dynamics get state significantly non-zero
mujoco.mj_forward(m, d)
mx = mjx.put_model(m)
dx = mjx.put_data(m, d)
# subtree velocity
mujoco.mj_subtreeVel(m, d)
dx = jax.jit(mjx.subtree_vel)(mx, dx)
_assert_attr_eq(d, dx, 'subtree_linvel')
_assert_attr_eq(d, dx, 'subtree_angmom')
if __name__ == '__main__':
absltest.main()