From b42780a5f4726333d815f28be086589637b75a06 Mon Sep 17 00:00:00 2001 From: Taylor Howell Date: Wed, 21 Aug 2024 16:20:39 -0700 Subject: [PATCH] Add subtree_vel function to MJX. This function matches mj_subtreeVel. PiperOrigin-RevId: 666078305 Change-Id: Ia0ecbf8279e914f49ee60b862541d697e8d4abb1 --- mjx/mujoco/mjx/__init__.py | 1 + mjx/mujoco/mjx/_src/smooth.py | 93 ++++++++++++++++++++++++++++++ mjx/mujoco/mjx/_src/smooth_test.py | 19 ++++++ 3 files changed, 113 insertions(+) diff --git a/mjx/mujoco/mjx/__init__.py b/mjx/mujoco/mjx/__init__.py index a4a8a5da..d272dd13 100644 --- a/mjx/mujoco/mjx/__init__.py +++ b/mjx/mujoco/mjx/__init__.py @@ -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 diff --git a/mjx/mujoco/mjx/_src/smooth.py b/mjx/mujoco/mjx/_src/smooth.py index ed6f3dce..96918c6a 100644 --- a/mjx/mujoco/mjx/_src/smooth.py +++ b/mjx/mujoco/mjx/_src/smooth.py @@ -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 diff --git a/mjx/mujoco/mjx/_src/smooth_test.py b/mjx/mujoco/mjx/_src/smooth_test.py index f0efc07f..c8e5f162 100644 --- a/mjx/mujoco/mjx/_src/smooth_test.py +++ b/mjx/mujoco/mjx/_src/smooth_test.py @@ -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()