diff --git a/doc/changelog.rst b/doc/changelog.rst
index 1447bd7b..c175ccd3 100644
--- a/doc/changelog.rst
+++ b/doc/changelog.rst
@@ -16,8 +16,8 @@ General
MJX
^^^
-
- Added ``mocap_pos`` and ``mocap_quat`` in kinematics.
+- Added support for :ref:`spatial tendons ` with external sphere and cylinder wrapping.
Bug fixes
^^^^^^^^^
diff --git a/mjx/mujoco/mjx/_src/io.py b/mjx/mujoco/mjx/_src/io.py
index e5ae79c9..19f92b55 100644
--- a/mjx/mujoco/mjx/_src/io.py
+++ b/mjx/mujoco/mjx/_src/io.py
@@ -118,6 +118,30 @@ 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.'
+ )
+
for enum_field, enum_type, mj_type in (
(m.actuator_biastype, types.BiasType, mujoco.mjtBias),
(m.actuator_dyntype, types.DynType, mujoco.mjtDyn),
diff --git a/mjx/mujoco/mjx/_src/io_test.py b/mjx/mujoco/mjx/_src/io_test.py
index c848ba9c..aaaa3415 100644
--- a/mjx/mujoco/mjx/_src/io_test.py
+++ b/mjx/mujoco/mjx/_src/io_test.py
@@ -175,6 +175,7 @@ class ModelIOTest(parameterized.TestCase):
+
@@ -186,7 +187,7 @@ class ModelIOTest(parameterized.TestCase):
-
+
diff --git a/mjx/mujoco/mjx/_src/smooth.py b/mjx/mujoco/mjx/_src/smooth.py
index 7e70c880..27b32510 100644
--- a/mjx/mujoco/mjx/_src/smooth.py
+++ b/mjx/mujoco/mjx/_src/smooth.py
@@ -720,15 +720,14 @@ def tendon(m: Model, d: Data) -> Data:
# find consecutive sites, skipping tendon transitions
(pair_id,) = np.nonzero(np.diff(wrap_id_site) == 1)
wrap_id_site_pair = np.setdiff1d(wrap_id_site[pair_id], m.tendon_adr[1:] - 1)
- (tendon_id_site,) = np.nonzero(np.isin(m.tendon_adr, wrap_id_site_pair))
- id0 = m.wrap_objid[wrap_id_site_pair]
- id1 = m.wrap_objid[wrap_id_site_pair + 1]
+ wrap_objid_site0 = m.wrap_objid[wrap_id_site_pair]
+ wrap_objid_site1 = m.wrap_objid[wrap_id_site_pair + 1]
@jax.vmap
def _length_moment(pnt0, pnt1, body0, body1):
dif = pnt1 - pnt0
- length = jp.linalg.norm(dif)
+ length = math.norm(dif)
vec = jp.where(
length < mujoco.mjMINVAL, jp.array([1.0, 0.0, 0.0]), dif / length
)
@@ -741,51 +740,200 @@ def tendon(m: Model, d: Data) -> Data:
return length, moment
lengths_site, moments_site = _length_moment(
- d.site_xpos[id0], d.site_xpos[id1], m.site_bodyid[id0], m.site_bodyid[id1]
+ d.site_xpos[wrap_objid_site0],
+ d.site_xpos[wrap_objid_site1],
+ m.site_bodyid[wrap_objid_site0],
+ m.site_bodyid[wrap_objid_site1],
)
tendon_nsite = np.array([
sum((wrap_id_site_pair >= adr) & (wrap_id_site_pair < adr + num))
for adr, num in zip(m.tendon_adr, m.tendon_num)
])
- tendon_nsite = tendon_nsite[tendon_nsite > 0]
- tendon_wrapnum_site = tendon_nsite + 1
- tendon_with_site = sum([s > 0 for s in tendon_nsite])
+ tendon_wrapnum_site = np.array([
+ sum((wrap_id_site >= adr) & (wrap_id_site < adr + num))
+ for adr, num in zip(m.tendon_adr, m.tendon_num)
+ ])
+ tendon_has_site = tendon_nsite > 0
+ (tendon_id_site,) = np.nonzero(tendon_has_site)
+ tendon_nsite = tendon_nsite[tendon_has_site]
+ tendon_with_site = tendon_nsite.size
ten_site_id = np.repeat(np.arange(tendon_with_site), tendon_nsite)
length_site = jax.ops.segment_sum(lengths_site, ten_site_id, tendon_with_site)
moment_site = jax.ops.segment_sum(moments_site, ten_site_id, tendon_with_site)
- # assemble length and moment
- ten_length = (
- jp.zeros_like(d.ten_length)
- .at[np.concatenate([tendon_id_jnt, tendon_id_site])]
- .set(jp.concatenate([length_jnt, length_site]))
+ # process spatial sphere/cylinder wrap
+ (wrap_id_geom,) = np.nonzero(
+ (m.wrap_type == WrapType.SPHERE) | (m.wrap_type == WrapType.CYLINDER)
)
+
+ # get objid for site-geom-site instances
+ wrap_id_sitegeomsite = wrap_id_geom[:, None] + np.array([-1, 0, 1])[None]
+ wrap_objid_site0, wrap_objid_geom, wrap_objid_site1 = m.wrap_objid[
+ wrap_id_sitegeomsite
+ ].T
+
+ # get site positions before and after geom
+ site_pnt0 = d.site_xpos[wrap_objid_site0]
+ site_pnt1 = d.site_xpos[wrap_objid_site1]
+
+ # get geom information
+ geom_xpos = d.geom_xpos[wrap_objid_geom]
+ 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]
+
+ # get body ids for site-geom-site instances
+ body_id_site0 = m.site_bodyid[wrap_objid_site0]
+ body_id_geom = m.geom_bodyid[wrap_objid_geom]
+ body_id_site1 = m.site_bodyid[wrap_objid_site1]
+
+ # find wrap object sidesites (if any exist)
+ side_id = np.round(m.wrap_prm[wrap_id_geom]).astype(int)
+ side = d.site_xpos[side_id]
+ has_sidesite = np.expand_dims(np.array(side_id >= 0), -1)
+
+ # 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,
+ )
+ lengths_geomgeom = lengths_geomgeom.reshape(-1)
+
+ # identify geoms where wrap does not occur
+ no_geom_wrap = lengths_geomgeom < 0
+ wrap_objid_geom_skip = jp.where(no_geom_wrap, 0, wrap_objid_geom)
+
+ # compute lengths for site-site (no wrap), site-geom, and geom-site segments
+ def _distance(p0, p1):
+ return jax.vmap(lambda x, y: math.norm(x - y))(p0, p1)
+
+ lengths_sitesite = _distance(site_pnt0, site_pnt1)
+ lengths_sitegeom = _distance(site_pnt0, geom_pnt0)
+ lengths_geomsite = _distance(geom_pnt1, site_pnt1)
+
+ # select length segments according to geom wrap
+ lengths_geom = jp.where(
+ no_geom_wrap,
+ lengths_sitesite,
+ lengths_sitegeom + lengths_geomgeom + lengths_geomsite,
+ )
+
+ # compute moments for site-site (no wrap), site-geom, geom-geom, and geom-site
+ # segments
+ _, moments_sitesite = _length_moment(
+ site_pnt0, site_pnt1, body_id_site0, body_id_site1
+ )
+ _, moments_sitegeom = _length_moment(
+ site_pnt0, geom_pnt0, body_id_site0, body_id_geom
+ )
+ _, moments_geomgeom = _length_moment(
+ geom_pnt0, geom_pnt1, body_id_geom, body_id_geom
+ )
+ _, moments_geomsite = _length_moment(
+ geom_pnt1, site_pnt1, body_id_geom, body_id_site1
+ )
+
+ # select moment segments according to geom wrap
+ moments_geom = jp.where(
+ no_geom_wrap[:, None],
+ moments_sitesite,
+ moments_sitegeom + moments_geomgeom + moments_geomsite,
+ )
+
+ # construct number of site-geom-site instances per tendon
+ tendon_ngeom = np.array([
+ sum((wrap_id_geom >= adr) & (wrap_id_geom < adr + num))
+ for adr, num in zip(m.tendon_adr, m.tendon_num)
+ ])
+ tendon_has_geom = tendon_ngeom > 0
+ tendon_ngeom = tendon_ngeom[tendon_has_geom]
+
+ # identify tendons with at least one site-geom-site instance
+ (tendon_id_geom,) = np.nonzero(tendon_has_geom)
+
+ # combine site-geom-site segment lengths and moments for each tendon
+ tendon_with_geom = tendon_ngeom.size
+ ten_geom_id = np.repeat(np.arange(tendon_with_geom), tendon_ngeom)
+
+ length_geom = jax.ops.segment_sum(lengths_geom, ten_geom_id, tendon_with_geom)
+ moment_geom = jax.ops.segment_sum(moments_geom, ten_geom_id, tendon_with_geom)
+
+ # calculate number of wrap objects per tendon, based on geom wrap
+ wrapnums_geom = jp.where(no_geom_wrap, 0, 2)
+ tendon_wrapnum_geom = jax.ops.segment_sum(
+ wrapnums_geom, ten_geom_id, tendon_with_geom
+ )
+
+ # assemble length and moment
+ ten_length = jp.zeros_like(d.ten_length).at[tendon_id_jnt].set(length_jnt)
+ ten_length = ten_length.at[tendon_id_site].add(length_site)
+ ten_length = ten_length.at[tendon_id_geom].add(length_geom)
+
ten_moment = (
jp.zeros_like(d.ten_J)
.at[adr_moment_jnt, dofadr_moment_jnt]
.set(moment_jnt)
)
- ten_moment = ten_moment.at[tendon_id_site].set(moment_site)
+ ten_moment = ten_moment.at[tendon_id_site].add(moment_site)
+ ten_moment = ten_moment.at[tendon_id_geom].add(moment_geom)
- # wrap
- wrap_xpos = jp.concatenate([
- d.site_xpos[m.wrap_objid[wrap_id_site]],
- jp.zeros((2 * m.nwrap - nwrap_site, 3)),
- ]).reshape((m.nwrap, 6))
+ # construct wrap addresses
+ wrap_adr_site = []
+ wrap_adr_geom = []
- ten_wrapnum = np.zeros(m.ntendon)
- ten_wrapnum[tendon_id_site] = tendon_wrapnum_site
+ count = 0
+ for wrap_type in m.wrap_type:
+ if wrap_type == WrapType.SITE:
+ wrap_adr_site.append(count)
+ count += 1
+ elif wrap_type in (WrapType.SPHERE, WrapType.CYLINDER):
+ wrap_adr_geom.append(count)
+ wrap_adr_geom.append(count + 1)
+ count += 2
- ten_wrapadr = [0]
- for wn in ten_wrapnum[:-1]:
- ten_wrapadr.append(ten_wrapadr[-1] + wn)
- ten_wrapadr = np.array(ten_wrapadr)
+ wrap_adr_site = np.array(wrap_adr_site).astype(int)
+ wrap_adr_geom = np.array(wrap_adr_geom).astype(int)
+ wrap_adr_sitegeom = np.concatenate([wrap_adr_site, wrap_adr_geom])
- wrap_obj = np.zeros(m.nwrap * 2, dtype=int)
- wrap_obj[:nwrap_site] = -1
- wrap_obj = wrap_obj.reshape((-1, 2))
+ ten_wrapnum = jp.array(tendon_wrapnum_site)
+ ten_wrapnum = ten_wrapnum.at[tendon_id_geom].add(tendon_wrapnum_geom)
+
+ ten_wrapadr = jp.concatenate([jp.array([0]), jp.cumsum(ten_wrapnum)[:-1]])
+
+ xpos_site = d.site_xpos[m.wrap_objid[wrap_id_site]]
+ xpos_geom = jp.hstack([geom_pnt0, geom_pnt1]).reshape((-1, 3))
+
+ # sort objects, moving no wrap geoms to bottom rows
+ nwrap_sitegeom = wrap_adr_sitegeom.size
+ wrap_adr_sitegeom_sort = np.argsort(wrap_adr_sitegeom)
+
+ skipped = (
+ jp.zeros(count, dtype=bool)
+ .at[wrap_adr_geom]
+ .set(jp.repeat(no_geom_wrap, 2).reshape(-1))
+ )
+ sort = jp.argsort(skipped)
+
+ wrap_xpos = jp.concatenate([xpos_site, xpos_geom])[wrap_adr_sitegeom_sort]
+ wrap_xpos = jp.concatenate(
+ [wrap_xpos[sort], jp.zeros((2 * m.nwrap - nwrap_sitegeom, 3))]
+ ).reshape((m.nwrap, 6))
+
+ wrap_obj = jp.concatenate([
+ -1 * jp.ones(nwrap_site, dtype=int),
+ jp.repeat(wrap_objid_geom_skip, 2).reshape(-1),
+ ])[wrap_adr_sitegeom_sort]
+ wrap_obj = jp.concatenate(
+ [wrap_obj[sort], jp.zeros(2 * m.nwrap - nwrap_sitegeom, dtype=int)]
+ ).reshape((m.nwrap, 2))
return d.replace(
ten_length=ten_length,
@@ -793,7 +941,7 @@ def tendon(m: Model, d: Data) -> Data:
ten_wrapadr=jp.array(ten_wrapadr, dtype=int),
ten_wrapnum=jp.array(ten_wrapnum, dtype=int),
wrap_xpos=wrap_xpos,
- wrap_obj=jp.array(wrap_obj),
+ wrap_obj=jp.array(wrap_obj, dtype=int),
)
diff --git a/mjx/mujoco/mjx/_src/smooth_test.py b/mjx/mujoco/mjx/_src/smooth_test.py
index b3dcfe97..9892e740 100644
--- a/mjx/mujoco/mjx/_src/smooth_test.py
+++ b/mjx/mujoco/mjx/_src/smooth_test.py
@@ -238,9 +238,12 @@ class TendonTest(parameterized.TestCase):
@parameterized.parameters(
'tendon/fixed.xml',
- 'tendon/site.xml',
'tendon/fixed_site.xml',
+ 'tendon/fixed_site_wrap.xml',
'tendon/no_tendon.xml',
+ 'tendon/site.xml',
+ 'tendon/site_wrap.xml',
+ 'tendon/wrap_sidesite.xml',
)
def test_tendon(self, filename):
"""Tests MJX tendon function matches MuJoCo mj_tendon."""
@@ -259,10 +262,8 @@ class TendonTest(parameterized.TestCase):
_assert_eq(d.ten_J, dx.ten_J, 'ten_J')
_assert_eq(d.ten_wrapnum, dx.ten_wrapnum, 'ten_wrapnum')
_assert_eq(d.ten_wrapadr, dx.ten_wrapadr, 'ten_wrapadr')
- if d.wrap_obj.shape == dx.wrap_obj.shape:
- _assert_eq(d.wrap_obj, dx.wrap_obj, 'wrap_obj')
- if d.wrap_xpos.shape == dx.wrap_xpos.shape:
- _assert_eq(d.wrap_xpos, dx.wrap_xpos, 'wrap_xpos')
+ _assert_eq(d.wrap_obj, dx.wrap_obj, 'wrap_obj')
+ _assert_eq(d.wrap_xpos, dx.wrap_xpos, 'wrap_xpos')
if __name__ == '__main__':
diff --git a/mjx/mujoco/mjx/_src/support.py b/mjx/mujoco/mjx/_src/support.py
index 891c973c..6345672e 100644
--- a/mjx/mujoco/mjx/_src/support.py
+++ b/mjx/mujoco/mjx/_src/support.py
@@ -346,3 +346,214 @@ def contact_force_dim(
raise NotImplementedError('Elliptic cone force is not implemented yet.')
else:
raise ValueError(f'Unknown cone type: {m.opt.cone}.')
+
+
+def length_circle(
+ p0: jax.Array, p1: jax.Array, ind: jax.Array, rad: jax.Array
+) -> jax.Array:
+ """Compute length of circle."""
+ # compute angle between 0 and pi
+ p0n = math.normalize(p0).reshape(-1)
+ p1n = math.normalize(p1).reshape(-1)
+
+ angle = jp.arccos(jp.dot(p0n, p1n))
+
+ # flip if necessary
+ cross = p0[1] * p1[0] - p0[0] * p1[1]
+ flip = ((cross > 0) & (ind != 0)) | ((cross < 0) & (ind == 0))
+ angle = jp.where(flip, 2 * jp.pi - angle, angle)
+
+ return rad * angle
+
+
+def is_intersect(
+ p1: jax.Array, p2: jax.Array, p3: jax.Array, p4: jax.Array
+) -> jax.Array:
+ """Check for intersection between two lines defined by their endpoints."""
+ # compute determinant
+ det = (p4[1] - p3[1]) * (p2[0] - p1[0]) - (p4[0] - p3[0]) * (p2[1] - p1[1])
+
+ # compute intersection point on each line
+ a = (
+ (p4[0] - p3[0]) * (p1[1] - p3[1]) - (p4[1] - p3[1]) * (p1[0] - p3[0])
+ ) / det
+ b = (
+ (p2[0] - p1[0]) * (p1[1] - p3[1]) - (p2[1] - p1[1]) * (p1[0] - p3[0])
+ ) / det
+
+ return (a >= 0) & (a <= 1) & (b >= 0) & (b <= 1)
+
+
+def wrap_circle(
+ d: jax.Array, sd: jax.Array, sidesite: jax.Array, rad: jax.Array
+) -> Tuple[jax.Array, jax.Array]:
+ """Compute circle wrap arc length and end points."""
+ # check cases
+ sqlen0 = d[0] ** 2 + d[1] ** 2
+ sqlen1 = d[2] ** 2 + d[3] ** 2
+ sqrad = rad * rad
+ dif = jp.array([d[2] - d[0], d[3] - d[1]])
+ dd = dif[0] ** 2 + dif[1] ** 2
+ a = jp.clip(-(dif[0] * d[0] + dif[1] * d[1]) / dd, 0, 1)
+ seg = jp.array([a * dif[0] + d[0], a * dif[1] + d[1]])
+
+ point_inside0 = sqlen0 < sqrad
+ point_inside1 = sqlen1 < sqrad
+ circle_too_small = rad < mujoco.mjMINVAL
+ points_too_close = dd < mujoco.mjMINVAL
+
+ intersect_and_side = (seg[0] ** 2 + seg[1] ** 2 > sqrad) & (
+ jp.where(sidesite, 0, 1) | (jp.dot(sd, seg) >= 0)
+ )
+
+ # construct the two solutions, compute goodness
+ def _sol(sgn):
+ sqrt0 = jp.sqrt(sqlen0 - sqrad)
+ sqrt1 = jp.sqrt(sqlen1 - sqrad)
+
+ d00 = (d[0] * sqrad + sgn * rad * d[1] * sqrt0) / sqlen0
+ d01 = (d[1] * sqrad - sgn * rad * d[0] * sqrt0) / sqlen0
+ d10 = (d[2] * sqrad - sgn * rad * d[3] * sqrt1) / sqlen1
+ d11 = (d[3] * sqrad + sgn * rad * d[2] * sqrt1) / sqlen1
+
+ sol = jp.array([[d00, d01], [d10, d11]])
+
+ # goodness: close to sd, or shorter path
+ tmp0 = sol[0] + sol[1]
+ tmp0 = math.normalize(tmp0).reshape(-1)
+ good0 = jp.dot(tmp0, sd)
+
+ tmp1 = (sol[0] - sol[1]).reshape(-1)
+ good1 = -jp.dot(tmp1, tmp1)
+
+ good = jp.where(sidesite, good0, good1)
+
+ # penalize for intersection
+ intersect = is_intersect(d[:2], sol[0], d[2:], sol[1])
+ good = jp.where(intersect, -10000, good)
+
+ return sol, good
+
+ sol, good = jax.vmap(_sol)(jp.array([1, -1]))
+
+ # select the better solution
+ i = jp.argmax(good)
+ sol = sol[i]
+ pnt = sol.reshape(-1)
+
+ # check for intersection
+ intersect = is_intersect(d[:2], pnt[:2], d[2:], pnt[2:])
+
+ # compute curve length
+ wlen = length_circle(sol[0], sol[1], i, rad)
+
+ # check cases
+ invalid = (
+ point_inside0
+ | point_inside1
+ | circle_too_small
+ | points_too_close
+ | intersect_and_side
+ | intersect
+ )
+
+ wlen = jp.where(invalid, -1, wlen)
+ pnt = jp.where(invalid, jp.zeros(4), pnt)
+
+ return wlen, pnt
+
+
+def wrap(
+ x0: jax.Array,
+ x1: jax.Array,
+ xpos: jax.Array,
+ xmat: jax.Array,
+ size: jax.Array,
+ side: jax.Array,
+ sidesite: jax.Array,
+ is_sphere: jax.Array,
+):
+ """Wrap tendon around sphere or cylinder."""
+ # map sites to wrap object's local frame
+ p0 = xmat.T @ (x0 - xpos)
+ p1 = xmat.T @ (x1 - xpos)
+
+ close_to_origin = (jp.linalg.norm(p0) < mujoco.mjMINVAL) | (
+ jp.linalg.norm(p1) < mujoco.mjMINVAL
+ )
+
+ # compute axes for sphere
+ # 1st axis
+ axis0 = p0
+ axis0 = math.normalize(axis0)
+
+ # compute normal to p0-0-p1 plane = cross(p0, p1)
+ normal = jp.cross(p0, p1)
+ normal, nrm = math.normalize_with_norm(normal)
+
+ # compute alternative normal (if (p0, p1) are parallel)
+ # find max component of axis0
+ axis_alt = jp.ones(3).at[jp.argmax(axis0)].set(0)
+ normal_alt = jp.cross(axis0, axis_alt)
+ normal_alt = math.normalize(normal_alt)
+
+ normal = jp.where(nrm < mujoco.mjMINVAL, normal_alt, normal)
+
+ # 2nd axis
+ axis1 = jp.cross(normal, axis0)
+ axis1 = math.normalize(axis1)
+
+ # set geom dependent axes
+ axis0 = jp.where(is_sphere, axis0, jp.array([1.0, 0.0, 0.0]))
+ axis1 = jp.where(is_sphere, axis1, jp.array([0.0, 1.0, 0.0]))
+
+ # project points in 2D frame: p => d
+ d = jp.array([
+ jp.dot(p0, axis0),
+ jp.dot(p0, axis1),
+ jp.dot(p1, axis0),
+ jp.dot(p1, axis1),
+ ])
+
+ # compute sidesite projection
+ s = xmat.T @ (side - xpos)
+ 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)
+ no_wrap = wlen < 0
+
+ # reconstruct 3D points in local frame: res
+ res0 = axis0 * pnt[0] + axis1 * pnt[1]
+ res1 = axis0 * pnt[2] + axis1 * pnt[3]
+ res = jp.concatenate([res0, res1])
+
+ # perform correction for cylinder case
+ l0 = jp.sqrt(
+ (p0[0] - res[0]) * (p0[0] - res[0]) + (p0[1] - res[1]) * (p0[1] - res[1])
+ )
+ l1 = jp.sqrt(
+ (p1[0] - res[3]) * (p1[0] - res[3]) + (p1[1] - res[4]) * (p1[1] - res[4])
+ )
+ r2 = p0[2] + (p1[2] - p0[2]) * l0 / (l0 + wlen + l1)
+ r5 = p0[2] + (p1[2] - p0[2]) * (l0 + wlen) / (l0 + wlen + l1)
+ height = jp.abs(r5 - r2)
+
+ wlen = jp.where(is_sphere, wlen, jp.sqrt(wlen * wlen + height * height))
+ res = jp.where(
+ is_sphere, res, res.at[jp.array([2, 5])].set(jp.concatenate([r2, r5]))
+ )
+
+ # map wrap points back to global frame
+ wpnt0 = xmat @ res[:3] + xpos
+ wpnt1 = xmat @ res[3:] + xpos
+
+ # check cases for no wrap
+ invalid = close_to_origin | no_wrap
+
+ wlen = jp.where(invalid, -1, wlen)
+ wpnt0 = jp.where(invalid, jp.zeros(3), wpnt0)
+ wpnt1 = jp.where(invalid, jp.zeros(3), wpnt1)
+
+ return wlen, wpnt0, wpnt1
diff --git a/mjx/mujoco/mjx/_src/types.py b/mjx/mujoco/mjx/_src/types.py
index 309b7338..a9ba040b 100644
--- a/mjx/mujoco/mjx/_src/types.py
+++ b/mjx/mujoco/mjx/_src/types.py
@@ -198,10 +198,14 @@ class WrapType(enum.IntEnum):
Members:
JOINT: constant moment arm
SITE: pass through site
+ SPHERE: wrap around sphere
+ CYLINDER: wrap around (infinite) cylinder
"""
JOINT = mujoco.mjtWrap.mjWRAP_JOINT
SITE = mujoco.mjtWrap.mjWRAP_SITE
- # unsupported: NONE, PULLEY, SPHERE, CYLINDER
+ SPHERE = mujoco.mjtWrap.mjWRAP_SPHERE
+ CYLINDER = mujoco.mjtWrap.mjWRAP_CYLINDER
+ # unsupported: NONE, PULLEY
class TrnType(enum.IntEnum):
diff --git a/mjx/mujoco/mjx/test_data/tendon/fixed_site_wrap.xml b/mjx/mujoco/mjx/test_data/tendon/fixed_site_wrap.xml
new file mode 100644
index 00000000..978e7fa3
--- /dev/null
+++ b/mjx/mujoco/mjx/test_data/tendon/fixed_site_wrap.xml
@@ -0,0 +1,58 @@
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
diff --git a/mjx/mujoco/mjx/test_data/tendon/site_wrap.xml b/mjx/mujoco/mjx/test_data/tendon/site_wrap.xml
new file mode 100644
index 00000000..1e0358b8
--- /dev/null
+++ b/mjx/mujoco/mjx/test_data/tendon/site_wrap.xml
@@ -0,0 +1,54 @@
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
diff --git a/mjx/mujoco/mjx/test_data/tendon/wrap_sidesite.xml b/mjx/mujoco/mjx/test_data/tendon/wrap_sidesite.xml
new file mode 100644
index 00000000..11092f8e
--- /dev/null
+++ b/mjx/mujoco/mjx/test_data/tendon/wrap_sidesite.xml
@@ -0,0 +1,52 @@
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+