Add external sphere and cylinder wrapping for spatial tendons to MJX.
PiperOrigin-RevId: 681808698 Change-Id: I37612e2c6b781b759f9009e1bc6ae10b4e861c8d
This commit is contained in:
committed by
Copybara-Service
parent
dfd1f8fefc
commit
c7c2edf758
+1
-1
@@ -16,8 +16,8 @@ General
|
||||
|
||||
MJX
|
||||
^^^
|
||||
|
||||
- Added ``mocap_pos`` and ``mocap_quat`` in kinematics.
|
||||
- Added support for :ref:`spatial tendons <tendon-spatial>` with external sphere and cylinder wrapping.
|
||||
|
||||
Bug fixes
|
||||
^^^^^^^^^
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -175,6 +175,7 @@ class ModelIOTest(parameterized.TestCase):
|
||||
<joint name="arm" axis="0 1 0"/>
|
||||
<geom name="shoulder" type="sphere" size=".05"/>
|
||||
<site name="arm" pos="-.1 0 .05"/>
|
||||
<site name="sidesite" pos="0 0 0"/>
|
||||
</body>
|
||||
<body name="slider" pos=".05 0 -.2">
|
||||
<joint name="slider" type="slide" damping="1"/>
|
||||
@@ -186,7 +187,7 @@ class ModelIOTest(parameterized.TestCase):
|
||||
<tendon>
|
||||
<spatial name="rope" range="0 .35">
|
||||
<site site="slider"/>
|
||||
<geom geom="shoulder"/>
|
||||
<geom geom="shoulder" sidesite="sidesite"/>
|
||||
<site site="arm"/>
|
||||
</spatial>
|
||||
</tendon>
|
||||
|
||||
+177
-29
@@ -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),
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -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__':
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -0,0 +1,58 @@
|
||||
<mujoco model="fixed_site_wrap">
|
||||
<option>
|
||||
<flag contact="disable" gravity="disable"/>
|
||||
</option>
|
||||
<worldbody>
|
||||
<body>
|
||||
<joint name="joint0" type="hinge" axis="0 1 0"/>
|
||||
<geom type="capsule" size="0.05 0.5" fromto="0 0 0 0.5 0 0"/>
|
||||
<site name="site0" pos="0.25 0 0.1" size="0.025"/>
|
||||
<geom name="sphere0" type="sphere" pos="0.5 0 0.125" size="0.085"/>
|
||||
<body pos="0.5 0 0">
|
||||
<joint name="joint1" type="hinge" axis="0 1 0"/>
|
||||
<geom type="capsule" size="0.05 0.5" fromto="0 0 0 0.5 0 0"/>
|
||||
<site name="site1" pos="0.25 0 0.1" size="0.025"/>
|
||||
<geom name="sphere1" type="sphere" pos="0.5 0 0.125" size="0.085"/>
|
||||
<body pos="0.5 0 0">
|
||||
<joint name="joint2" type="hinge" axis="0 1 0"/>
|
||||
<geom type="capsule" size="0.05 0.5" fromto="0 0 0 0.5 0 0"/>
|
||||
<site name="site2" pos="0.25 0 0.1" size="0.025"/>
|
||||
<geom name="cylinder2" type="cylinder" contype="0" conaffinity="0" pos="0.5 0 0.25" euler="90 0 0" size="0.085 0.1"/>
|
||||
<body pos="0.5 0 0">
|
||||
<joint name="joint3" type="hinge" axis="0 1 0"/>
|
||||
<geom type="capsule" size="0.05 0.5" fromto="0 0 0 0.5 0 0"/>
|
||||
<site name="site3" pos="0.25 0 0.1" size="0.025"/>
|
||||
</body>
|
||||
</body>
|
||||
</body>
|
||||
</body>
|
||||
</worldbody>
|
||||
|
||||
<tendon>
|
||||
<spatial width="0.0125">
|
||||
<site site="site2"/>
|
||||
<site site="site3"/>
|
||||
</spatial>
|
||||
<fixed>
|
||||
<joint joint="joint0" coef=".1"/>
|
||||
<joint joint="joint1" coef=".2"/>
|
||||
<joint joint="joint2" coef=".3"/>
|
||||
</fixed>
|
||||
<spatial width="0.0125">
|
||||
<site site="site0"/>
|
||||
<geom geom="sphere0"/>
|
||||
<site site="site1"/>
|
||||
<geom geom="sphere1"/>
|
||||
<site site="site2"/>
|
||||
<geom geom="cylinder2"/>
|
||||
<site site="site3"/>
|
||||
</spatial>
|
||||
<spatial name="tendon2" width="0.0125">
|
||||
<site site="site0"/>
|
||||
<geom geom="sphere0"/>
|
||||
<site site="site2"/>
|
||||
<geom geom="cylinder2"/>
|
||||
<site site="site3"/>
|
||||
</spatial>
|
||||
</tendon>
|
||||
</mujoco>
|
||||
@@ -0,0 +1,54 @@
|
||||
<mujoco model="site_wrap">
|
||||
<option>
|
||||
<flag contact="disable" gravity="disable"/>
|
||||
</option>
|
||||
<worldbody>
|
||||
<light pos="0 0 10"/>
|
||||
<body>
|
||||
<joint name="joint0" type="hinge" axis="0 1 0"/>
|
||||
<geom type="capsule" size="0.05 0.5" fromto="0 0 0 0.5 0 0"/>
|
||||
<site name="site0" pos="0.25 0 0.1" size="0.025"/>
|
||||
<geom name="sphere0" type="sphere" pos="0.5 0 0.125" size="0.085"/>
|
||||
<body pos="0.5 0 0">
|
||||
<joint name="joint1" type="hinge" axis="0 1 0"/>
|
||||
<geom type="capsule" size="0.05 0.5" fromto="0 0 0 0.5 0 0"/>
|
||||
<site name="site1" pos="0.25 0 0.1" size="0.025"/>
|
||||
<geom name="sphere1" type="sphere" pos="0.5 0 0.125" size="0.085"/>
|
||||
<body pos="0.5 0 0">
|
||||
<joint name="joint2" type="hinge" axis="0 1 0"/>
|
||||
<geom type="capsule" size="0.05 0.5" fromto="0 0 0 0.5 0 0"/>
|
||||
<site name="site2" pos="0.25 0 0.1" size="0.025"/>
|
||||
<geom name="cylinder2" type="cylinder" contype="0" conaffinity="0" pos="0.5 0 0.25" euler="90 0 0" size="0.085 0.1"/>
|
||||
<body pos="0.5 0 0">
|
||||
<joint name="joint3" type="hinge" axis="0 1 0"/>
|
||||
<geom type="capsule" size="0.05 0.5" fromto="0 0 0 0.5 0 0"/>
|
||||
<site name="site3" pos="0.25 0 0.1" size="0.025"/>
|
||||
</body>
|
||||
</body>
|
||||
</body>
|
||||
</body>
|
||||
</worldbody>
|
||||
|
||||
<tendon>
|
||||
<spatial width="0.0125">
|
||||
<site site="site0"/>
|
||||
<geom geom="sphere0"/>
|
||||
<site site="site1"/>
|
||||
<geom geom="sphere1"/>
|
||||
<site site="site2"/>
|
||||
<geom geom="cylinder2"/>
|
||||
<site site="site3"/>
|
||||
</spatial>
|
||||
<spatial width="0.0125">
|
||||
<site site="site2"/>
|
||||
<site site="site3"/>
|
||||
</spatial>
|
||||
<spatial width="0.0125">
|
||||
<site site="site0"/>
|
||||
<geom geom="sphere0"/>
|
||||
<site site="site2"/>
|
||||
<geom geom="cylinder2"/>
|
||||
<site site="site3"/>
|
||||
</spatial>
|
||||
</tendon>
|
||||
</mujoco>
|
||||
@@ -0,0 +1,52 @@
|
||||
<mujoco model="wrap_sidesite">
|
||||
<option>
|
||||
<flag contact="disable" gravity="disable"/>
|
||||
</option>
|
||||
<worldbody>
|
||||
<light pos="0 0 10"/>
|
||||
<body>
|
||||
<joint name="joint0" type="hinge" axis="0 1 0"/>
|
||||
<geom type="capsule" size="0.05 0.5" fromto="0 0 0 0.5 0 0"/>
|
||||
<site name="site0" pos="0.25 0 0.1" size="0.025"/>
|
||||
<geom name="sphere0" type="sphere" pos="0.5 0 0.125" size="0.085"/>
|
||||
<site name="sidesite0" pos="0.5 0 0.25"/>
|
||||
<body pos="0.5 0 0">
|
||||
<joint name="joint1" type="hinge" axis="0 1 0"/>
|
||||
<geom type="capsule" size="0.05 0.5" fromto="0 0 0 0.5 0 0"/>
|
||||
<site name="site1" pos="0.25 0 0.1" size="0.025"/>
|
||||
<geom name="sphere1" type="sphere" pos="0.5 0 0.125" size="0.085"/>
|
||||
<body pos="0.5 0 0">
|
||||
<joint name="joint2" type="hinge" axis="0 1 0"/>
|
||||
<geom type="capsule" size="0.05 0.5" fromto="0 0 0 0.5 0 0"/>
|
||||
<site name="site2" pos="0.25 0 0.1" size="0.025"/>
|
||||
<geom name="cylinder2" type="cylinder" contype="0" conaffinity="0" pos="0.5 0 0.25" euler="90 0 0" size="0.085 0.1"/>
|
||||
<site name="sidesite2" pos="0.5 0 0.5"/>
|
||||
<body pos="0.5 0 0">
|
||||
<joint name="joint3" type="hinge" axis="0 1 0"/>
|
||||
<geom type="capsule" size="0.05 0.5" fromto="0 0 0 0.5 0 0"/>
|
||||
<site name="site3" pos="0.25 0 0.1" size="0.025"/>
|
||||
</body>
|
||||
</body>
|
||||
</body>
|
||||
</body>
|
||||
</worldbody>
|
||||
|
||||
<tendon>
|
||||
<spatial width="0.0125">
|
||||
<site site="site0"/>
|
||||
<geom geom="sphere0" sidesite="sidesite0"/>
|
||||
<site site="site1"/>
|
||||
<geom geom="sphere1"/>
|
||||
<site site="site2"/>
|
||||
<geom geom="cylinder2" sidesite="sidesite2"/>
|
||||
<site site="site3"/>
|
||||
</spatial>
|
||||
<spatial width="0.0125">
|
||||
<site site="site0"/>
|
||||
<geom geom="sphere0"/>
|
||||
<site site="site2"/>
|
||||
<geom geom="cylinder2" sidesite="sidesite2"/>
|
||||
<site site="site3"/>
|
||||
</spatial>
|
||||
</tendon>
|
||||
</mujoco>
|
||||
Reference in New Issue
Block a user