From f4096bcab9baf48a330f66f1c172745400d12b10 Mon Sep 17 00:00:00 2001 From: Taylor Howell Date: Tue, 28 Jan 2025 03:20:18 -0800 Subject: [PATCH] Add sphere and cylinder internal wrapping for spatial tendons to MJX. PiperOrigin-RevId: 720507688 Change-Id: I49fda2839e68623ff5e1f3c46bdc3d77acfc36ef --- doc/changelog.rst | 4 + doc/mjx.rst | 8 +- mjx/mujoco/mjx/_src/io.py | 47 +++--- mjx/mujoco/mjx/_src/io_test.py | 42 ++--- mjx/mujoco/mjx/_src/smooth.py | 54 +++++-- mjx/mujoco/mjx/_src/support.py | 145 +++++++++++++++++- mjx/mujoco/mjx/_src/support_test.py | 121 +++++++++++++++ mjx/mujoco/mjx/_src/types.py | 8 + .../mjx/test_data/tendon/wrap_sidesite.xml | 15 +- 9 files changed, 373 insertions(+), 71 deletions(-) diff --git a/doc/changelog.rst b/doc/changelog.rst index 00c92e4d..7c9efba4 100644 --- a/doc/changelog.rst +++ b/doc/changelog.rst @@ -13,6 +13,10 @@ General attaching can be restored by setting the deep copy flag to 1. - Added :ref:`potential` and :ref:`kinetic` energy sensors. +MJX +^^^ +- Added support for spatial tendons with internal sphere and cylinder wrapping. + Version 3.2.7 (Jan 14, 2025) ---------------------------- diff --git a/doc/mjx.rst b/doc/mjx.rst index 55f9aab3..b8d882ae 100644 --- a/doc/mjx.rst +++ b/doc/mjx.rst @@ -220,7 +220,7 @@ The following features are **fully supported** in MJX: * - :ref:`Actuator Bias ` - ``NONE``, ``AFFINE``, ``MUSCLE`` * - :ref:`Tendon Wrapping ` - - ``JOINT``, ``SITE``, ``PULLEY`` + - ``JOINT``, ``SITE``, ``PULLEY``, ``SPHERE``, ``CYLINDER`` * - :ref:`Geom ` - ``PLANE``, ``HFIELD``, ``SPHERE``, ``CAPSULE``, ``BOX``, ``MESH`` are fully implemented. ``ELLIPSOID`` and ``CYLINDER`` are implemented but only collide with other primitives, note that ``BOX`` is implemented as a mesh. @@ -239,7 +239,7 @@ The following features are **fully supported** in MJX: * - Fluid Model - :ref:`flInertia` * - :ref:`Tendons ` - - :ref:`Fixed ` + - :ref:`Fixed `, :ref:`Spatial ` * - :ref:`Sensors ` - ``MAGNETOMETER``, ``CAMPROJECTION``, ``RANGEFINDER``, ``JOINTPOS``, ``TENDONPOS``, ``ACTUATORPOS``, ``BALLQUAT``, ``FRAMEPOS``, ``FRAMEXAXIS``, ``FRAMEYAXIS``, ``FRAMEZAXIS``, ``FRAMEQUAT``, ``SUBTREECOM``, ``CLOCK``, @@ -265,12 +265,8 @@ The following features are **in development** and coming soon: - ``IMPLICIT`` * - Dynamics - :ref:`Inverse ` - * - :ref:`Tendon Wrapping ` - - ``SPHERE``, ``CYLINDER`` (external wrapping is supported) * - Fluid Model - :ref:`flEllipsoid` - * - :ref:`Tendons ` - - :ref:`Spatial ` * - :ref:`Sensors ` - All except ``PLUGIN``, ``USER`` * - Lights diff --git a/mjx/mujoco/mjx/_src/io.py b/mjx/mujoco/mjx/_src/io.py index f4e5aaeb..89ed3804 100644 --- a/mjx/mujoco/mjx/_src/io.py +++ b/mjx/mujoco/mjx/_src/io.py @@ -123,29 +123,6 @@ def put_model( if t == mujoco.mjtGeom.mjGEOM_MESH: mesh_geomid.add(g) - # check for spatial tendon internal geom wrapping - if m.ntendon: - # find sphere or cylinder geoms (if any exist) - (wrap_id_geom,) = np.nonzero( - (m.wrap_type == mujoco.mjtWrap.mjWRAP_SPHERE) - | (m.wrap_type == mujoco.mjtWrap.mjWRAP_CYLINDER) - ) - wrap_objid_geom = m.wrap_objid[wrap_id_geom] - geom_pos = m.geom_pos[wrap_objid_geom] - geom_size = m.geom_size[wrap_objid_geom, 0] - - # find sidesites (if any exist) - side_id = np.round(m.wrap_prm[wrap_id_geom]).astype(int) - side = m.site_pos[side_id] - - # check for sidesite inside geom - if np.any( - (np.linalg.norm(side - geom_pos, axis=1) < geom_size) & (side_id >= 0) - ): - raise NotImplementedError( - 'Internal wrapping with sphere and cylinder geoms is not' - ' implemented for spatial tendons.' - ) # check for unsupported sensor and equality constraint combinations sensor_rne_postconstraint = ( @@ -199,6 +176,30 @@ def put_model( fields['opt'] = _make_option(m.opt, _full_compat=_full_compat) fields['stat'] = _make_statistic(m.stat) + # spatial tendon wrap inside + fields['wrap_inside_maxiter'] = 5 + fields['wrap_inside_tolerance'] = 1.0e-4 + fields['wrap_inside_z_init'] = 1.0 - 1.0e-5 + fields['is_wrap_inside'] = np.zeros(0, dtype=bool) + if m.nsite: + # find sphere or cylinder geoms (if any exist) + (wrap_id_geom,) = np.nonzero( + (m.wrap_type == mujoco.mjtWrap.mjWRAP_SPHERE) + | (m.wrap_type == mujoco.mjtWrap.mjWRAP_CYLINDER) + ) + wrap_objid_geom = m.wrap_objid[wrap_id_geom] + geom_pos = m.geom_pos[wrap_objid_geom] + geom_size = m.geom_size[wrap_objid_geom, 0] + + # find sidesites (if any exist) + side_id = np.round(m.wrap_prm[wrap_id_geom]).astype(int) + side = m.site_pos[side_id] + + # wrap inside flag + fields['is_wrap_inside'] = np.array( + (np.linalg.norm(side - geom_pos, axis=1) < geom_size) & (side_id >= 0) + ) + # Pre-compile meshes for MJX collisions. fields['mesh_convex'] = [None] * m.nmesh if not _full_compat: diff --git a/mjx/mujoco/mjx/_src/io_test.py b/mjx/mujoco/mjx/_src/io_test.py index cb9a9f75..01977792 100644 --- a/mjx/mujoco/mjx/_src/io_test.py +++ b/mjx/mujoco/mjx/_src/io_test.py @@ -20,6 +20,7 @@ import jax from jax import numpy as jp import mujoco from mujoco import mjx +from mujoco.mjx._src import test_util # pylint: disable=g-importing-member from mujoco.mjx._src.types import ConeType # pylint: enable=g-importing-member @@ -171,33 +172,6 @@ class ModelIOTest(parameterized.TestCase): ) ) - def test_spatial_tendon_not_implemented(self): - with self.assertRaises(NotImplementedError): - mjx.put_model(mujoco.MjModel.from_xml_string(""" - - - - - - - - - - - - - - - - - - - - - - - """)) - def test_margin_gap_mesh_not_implemented(self): with self.assertRaises(NotImplementedError): mjx.put_model(mujoco.MjModel.from_xml_string(""" @@ -225,6 +199,20 @@ class ModelIOTest(parameterized.TestCase): """)) + def test_wrap_inside(self): + m = test_util.load_test_file('tendon/wrap_sidesite.xml') + mx0 = mjx.put_model(m) + np.testing.assert_equal( + mx0.is_wrap_inside, + np.array([1, 0, 1, 0, 1, 1, 0]), + ) + m.site_pos[2] = m.site_pos[1] + mx1 = mjx.put_model(m) + np.testing.assert_equal( + mx1.is_wrap_inside, + np.array([0, 0, 1, 0, 1, 0, 0]), + ) + class DataIOTest(parameterized.TestCase): """IO tests for mjx.Data.""" diff --git a/mjx/mujoco/mjx/_src/smooth.py b/mjx/mujoco/mjx/_src/smooth.py index f52391ed..87750c04 100644 --- a/mjx/mujoco/mjx/_src/smooth.py +++ b/mjx/mujoco/mjx/_src/smooth.py @@ -812,6 +812,7 @@ def tendon(m: Model, d: Data) -> Data: geom_xmat = d.geom_xmat[wrap_objid_geom] geom_size = m.geom_size[wrap_objid_geom, 0] geom_type = m.wrap_type[wrap_id_geom] + is_sphere = geom_type == WrapType.SPHERE # get body ids for site-geom-site instances body_id_site0 = m.site_bodyid[wrap_objid_site0] @@ -823,17 +824,52 @@ def tendon(m: Model, d: Data) -> Data: side = d.site_xpos[side_id] has_sidesite = np.expand_dims(np.array(side_id >= 0), -1) + # wrap inside + # TODO(taylorhowell): check that is_wrap_inside is consistent with + # site and geom relative positions + (wrap_inside_id,) = np.nonzero(m.is_wrap_inside) + (wrap_outside_id,) = np.nonzero(~m.is_wrap_inside) + # compute geom wrap length and connect points (if wrap occurs) - lengths_geomgeom, geom_pnt0, geom_pnt1 = jax.vmap(support.wrap)( - site_pnt0, - site_pnt1, - geom_xpos, - geom_xmat, - geom_size, - side, - has_sidesite, - geom_type == WrapType.SPHERE, + v_wrap = jax.vmap( + support.wrap, in_axes=(0, 0, 0, 0, 0, 0, 0, 0, None, None, None, None) ) + lengths_inside, pnt0_inside, pnt1_inside = v_wrap( + site_pnt0[wrap_inside_id], + site_pnt1[wrap_inside_id], + geom_xpos[wrap_inside_id], + geom_xmat[wrap_inside_id], + geom_size[wrap_inside_id], + side[wrap_inside_id], + has_sidesite[wrap_inside_id], + is_sphere[wrap_inside_id], + True, + m.wrap_inside_maxiter, + m.wrap_inside_tolerance, + m.wrap_inside_z_init, + ) + + lengths_outside, pnt0_outside, pnt1_outside = v_wrap( + site_pnt0[wrap_outside_id], + site_pnt1[wrap_outside_id], + geom_xpos[wrap_outside_id], + geom_xmat[wrap_outside_id], + geom_size[wrap_outside_id], + side[wrap_outside_id], + has_sidesite[wrap_outside_id], + is_sphere[wrap_outside_id], + False, + m.wrap_inside_maxiter, + m.wrap_inside_tolerance, + m.wrap_inside_z_init, + ) + + wrap_id = np.argsort(np.concatenate([wrap_inside_id, wrap_outside_id])) + vstack_ = lambda x, y: jp.vstack([x, y])[wrap_id] + lengths_geomgeom = vstack_(lengths_inside, lengths_outside) + geom_pnt0 = vstack_(pnt0_inside, pnt0_outside) + geom_pnt1 = vstack_(pnt1_inside, pnt1_outside) + lengths_geomgeom = lengths_geomgeom.reshape(-1) # identify geoms where wrap does not occur diff --git a/mjx/mujoco/mjx/_src/support.py b/mjx/mujoco/mjx/_src/support.py index 01326d52..041942b9 100644 --- a/mjx/mujoco/mjx/_src/support.py +++ b/mjx/mujoco/mjx/_src/support.py @@ -692,6 +692,136 @@ def wrap_circle( return wlen, pnt +def wrap_inside( + end: jax.Array, + radius: jax.Array, + maxiter: int, + tolerance: float, + z_init: float, +) -> Tuple[jax.Array, jax.Array]: + """Compute 2D inside wrap point. + + Args: + end: 2D points + radius: radius of circle + + Returns: + status: 0 if wrap, else -1 + concatentated 2D wrap points: jax.Array + """ + mjMINVAL = mujoco.mjMINVAL # pylint: disable=invalid-name + + # constants + len0 = math.norm(end[:2]) + len1 = math.norm(end[2:]) + dif = jp.array([end[2] - end[0], end[3] - end[1]]) + dd = dif[0] * dif[0] + dif[1] * dif[1] + + # either point inside circle or circle too small: no wrap + no_wrap0 = ( + (len0 <= radius) + | (len1 <= radius) + | (radius < mjMINVAL) + | (len0 < mjMINVAL) + | (len1 < mjMINVAL) + ) + + # find nearest point on line segment to origin: d0 + a*dif + a = -1 * (dif[0] * end[0] + dif[1] * end[1]) / jp.maximum(mjMINVAL, dd) + tmp = end[:2] + a * dif + + # segment-circle intersection: no wrap + no_wrap1 = (dd > mjMINVAL) & (a > 0) & (a < 1) & (math.norm(tmp) <= radius) + + # prepare default in case of numerical failure: average + pnt_avg = 0.5 * jp.array([end[0] + end[2], end[1] + end[3]]) + pnt_avg = radius * math.normalize(pnt_avg) + + # compute function parameters: asin(A*z) + asin(B*z) - 2*asin(z) + G = 0 + A = radius / jp.maximum(mjMINVAL, len0) # pylint: disable=invalid-name + B = radius / jp.maximum(mjMINVAL, len1) # pylint: disable=invalid-name + cosG = (len0 * len0 + len1 * len1 - dd) / jp.maximum(mjMINVAL, 2 * len0 * len1) # pylint: disable=invalid-name + + no_wrap2 = cosG < -1 + mjMINVAL + early_return0 = cosG > 1 - mjMINVAL + + G = jp.arccos(cosG) # pylint: disable=invalid-name + + # initialize solver + z = jp.array([z_init]) + f = jp.arcsin(A * z) + jp.arcsin(B * z) - 2 * jp.arcsin(z) + G + + # make sure initialization is not on the other side + early_return1 = f > 0 + + # iteratively solve with Newton's method + def _newton(carry, _): + # unpack + z, f, status_prev = carry + + # check current solution + converged = jp.abs(f) <= tolerance + + # compute derivative + df = ( + A / jp.maximum(mjMINVAL, jp.sqrt(1 - z * z * A * A)) + + B / jp.maximum(mjMINVAL, jp.sqrt(1 - z * z * B * B)) + - 2 / jp.maximum(mjMINVAL, jp.sqrt(1 - z * z)) + ) + + # check sign; SHOULD NOT OCCUR + status0 = df > -mjMINVAL + + # new point + z_next = z - (1 - converged) * f / jp.where( + jp.abs(df) < mjMINVAL, mjMINVAL, df + ) + + # make sure we are moving to the left; SHOULD NOT OCCUR + status1 = z_next > z + + # evaluate solution + f_next = ( + jp.arcsin(A * z_next) + + jp.arcsin(B * z_next) + - 2 * jp.arcsin(z_next) + + G + ) + + # exit if positive; SHOULD NOT OCCUR + status2 = f_next > tolerance + + return ( + z_next, + f_next, + status_prev | status0 | status1 | status2, + ), None + + # TODO(taylorhowell): compare performance of jax.lax.scan and jax.lax.while_loop + z, _, early_return2 = jax.lax.scan( + _newton, (z, f, jp.array([False])), None, maxiter + )[0] + + # finalize: rotation by ang from vec = a or b, depending on cross(a,b) sign + sign = end[0] * end[3] - end[1] * end[2] > 0 + vec = jp.where(sign, end[:2], end[2:]) + vec = math.normalize(vec) + ang = jp.arcsin(z) - jp.where(sign, jp.arcsin(A * z), jp.arcsin(B * z)) + pnt_sol = radius * jp.array([ + jp.cos(ang) * vec[0] - jp.sin(ang) * vec[1], + jp.sin(ang) * vec[0] + jp.cos(ang) * vec[1], + ]).reshape(-1) + + no_wrap = no_wrap0 | no_wrap1 | no_wrap2 + early_return = early_return0 | early_return1 | early_return2 + status = -1 * no_wrap * jp.ones(1) + + pnt = jp.where(early_return, pnt_avg, pnt_sol) + pnt = jp.where(no_wrap, jp.zeros(2), pnt) + + return status, jp.concatenate([pnt, pnt]) + + def wrap( x0: jax.Array, x1: jax.Array, @@ -701,7 +831,11 @@ def wrap( side: jax.Array, sidesite: jax.Array, is_sphere: jax.Array, -): + is_wrap_inside: bool, + wrap_inside_maxiter: int, + wrap_inside_tolerance: float, + wrap_inside_z_init: float, +) -> Tuple[jax.Array, jax.Array, jax.Array]: """Wrap tendon around sphere or cylinder.""" # map sites to wrap object's local frame p0 = xmat.T @ (x0 - xpos) @@ -749,8 +883,13 @@ def wrap( sd = jp.array([jp.dot(s, axis0), jp.dot(s, axis1)]) sd = math.normalize(sd) * size - # TODO(taylorhowell): implement wrap_inside for internal wrapping case - wlen, pnt = wrap_circle(d, sd, sidesite, size) + if is_wrap_inside: + wlen, pnt = wrap_inside( + d, size, wrap_inside_maxiter, wrap_inside_tolerance, wrap_inside_z_init + ) + else: + wlen, pnt = wrap_circle(d, sd, sidesite, size) + no_wrap = wlen < 0 # reconstruct 3D points in local frame: res diff --git a/mjx/mujoco/mjx/_src/support_test.py b/mjx/mujoco/mjx/_src/support_test.py index eb4c5b20..6fadc843 100644 --- a/mjx/mujoco/mjx/_src/support_test.py +++ b/mjx/mujoco/mjx/_src/support_test.py @@ -344,6 +344,127 @@ class SupportTest(parameterized.TestCase): force = force.at[3:].set(dx.contact.frame[j] @ force[3:]) np.testing.assert_allclose(result, force, rtol=1e-5, atol=2) + def test_wrap_inside(self): + maxiter = 5 + tolerance = 1.0e-4 + z_init = 1.0 - 1.0e-5 + + # len0 <= radius + np.testing.assert_equal( + support.wrap_inside( + jp.array([1.0, 0, 0, 0]), + jp.array([1.0]), + maxiter, + tolerance, + z_init, + )[0], + jp.array([-1]), + ) + + # len1 <= radius + np.testing.assert_equal( + support.wrap_inside( + jp.array([0, 0, 1.0, 0]), + jp.array([1.0]), + maxiter, + tolerance, + z_init, + )[0], + jp.array([-1]), + ) + + # radius < mjMINVAL + np.testing.assert_equal( + support.wrap_inside( + jp.array([1, 0, 0, 1]), jp.array([0.1 * mujoco.mjMINVAL]), maxiter, + tolerance, + z_init, + )[0], + jp.array([-1]), + ) + + # len0 < mjMINVAL and radius < mjMINVAL + np.testing.assert_equal( + support.wrap_inside( + jp.array([0.1 * mujoco.mjMINVAL, 0, 0, 0]), + jp.array([0.1 * mujoco.mjMINVAL]), + maxiter, + tolerance, + z_init, + )[0], + jp.array([-1]), + ) + + # len1 < mjMINVAL and radius < mjMINVAL + np.testing.assert_equal( + support.wrap_inside( + jp.array([0, 0, 0.1 * mujoco.mjMINVAL, 0]), + jp.array([0.1 * mujoco.mjMINVAL]), + maxiter, + tolerance, + z_init, + )[0], + jp.array([-1]), + ) + + # wrap: p0 = [1, 0], p1 = [0, 1] + status, pnt = support.wrap_inside( + jp.array([1, 0, 0, 1]), jp.array([0.5]), maxiter, tolerance, z_init + ) + np.testing.assert_allclose( + pnt, + jp.array([0.353553, 0.353553, 0.353553, 0.353553]), + atol=1e-3, + rtol=1e-3, + ) + np.testing.assert_equal(status, jp.array([0])) + + # no wrap, point on circle: p0 = [1, 0], p1 = [0, 0.5] + status, pnt = support.wrap_inside( + jp.array([1, 0, 0, 0.5]), jp.array([0.5]), maxiter, tolerance, z_init + ) + np.testing.assert_allclose( + pnt, + jp.zeros(4), + atol=1e-3, + rtol=1e-3, + ) + np.testing.assert_equal(status, jp.array([-1])) + + # no wrap, segment-circle intersection: p0 = [0.75, 0], p1 = [0, 0.51] + status, pnt = support.wrap_inside( + jp.array([0.75, 0, 0, 0.51]), + jp.array([0.5]), + maxiter, + tolerance, + z_init, + ) + np.testing.assert_allclose( + pnt, + jp.zeros(4), + atol=1e-3, + rtol=1e-3, + ) + np.testing.assert_equal(status, jp.array([-1])) + + # wrap: p0 = [-0.5, 1], p1 = [0.5, 1] + status, pnt = support.wrap_inside( + jp.array([-0.5, 1, 0.5, 1]), + jp.array([0.5]), + maxiter, + tolerance, + z_init, + ) + np.testing.assert_allclose( + pnt, + jp.array([0, 0.5, 0, 0.5]), + atol=1e-3, + rtol=1e-3, + ) + np.testing.assert_equal(status, jp.array([0])) + + # TODO(taylorhowell): improve wrap_inside testing with additional test cases + def test_muscle_gain_length(self): lmin = 0.5 lmax = 1.5 diff --git a/mjx/mujoco/mjx/_src/types.py b/mjx/mujoco/mjx/_src/types.py index 030f4a27..19dd56b2 100644 --- a/mjx/mujoco/mjx/_src/types.py +++ b/mjx/mujoco/mjx/_src/types.py @@ -779,6 +779,10 @@ class Model(PyTreeNode): wrap_type: wrap object type (mjtWrap) (nwrap,) wrap_objid: object id: geom, site, joint (nwrap,) wrap_prm: divisor, joint coef, or site id (nwrap,) + wrap_inside_maxiter: maximum iterations for wrap_inside + wrap_inside_tolerance: tolerance for wrap_inside + wrap_inside_z_init: initialization for wrap_inside + is_wrap_inside: spatial tendon sidesite inside geom (nwrapinside,) actuator_trntype: transmission type (mjtTrn) (nu,) actuator_dyntype: dynamics type (mjtDyn) (nu,) actuator_gaintype: gain type (mjtGain) (nu,) @@ -1109,6 +1113,10 @@ class Model(PyTreeNode): wrap_type: np.ndarray wrap_objid: np.ndarray wrap_prm: np.ndarray + wrap_inside_maxiter: int = _restricted_to('mjx') + wrap_inside_tolerance: float = _restricted_to('mjx') + wrap_inside_z_init: float = _restricted_to('mjx') + is_wrap_inside: np.ndarray = _restricted_to('mjx') actuator_trntype: np.ndarray actuator_dyntype: np.ndarray actuator_gaintype: np.ndarray diff --git a/mjx/mujoco/mjx/test_data/tendon/wrap_sidesite.xml b/mjx/mujoco/mjx/test_data/tendon/wrap_sidesite.xml index 11092f8e..36bb758b 100644 --- a/mjx/mujoco/mjx/test_data/tendon/wrap_sidesite.xml +++ b/mjx/mujoco/mjx/test_data/tendon/wrap_sidesite.xml @@ -10,6 +10,7 @@ + @@ -21,6 +22,7 @@ + @@ -34,16 +36,23 @@ - + - + - + + + + + + + +