Add cylinder collisions using SDFs.
PiperOrigin-RevId: 626998709 Change-Id: I669a639b1cabf5ed522c73a07e54684d675fa396
This commit is contained in:
committed by
Copybara-Service
parent
9aa885aa63
commit
c2d0c5dd1b
+4
-3
@@ -31,12 +31,13 @@ MJX
|
||||
9. Changed the way meshes are organized on device to speed up collision detection when a mesh is replicated for many
|
||||
geoms.
|
||||
10. Fixed a bug where capsules might be ignored in broadphase colliision checking.
|
||||
11. Added cylinder collisions using SDFs.
|
||||
|
||||
Bug fixes
|
||||
^^^^^^^^^
|
||||
11. Defaults of lights were not being saved, now fixed.
|
||||
12. Prevent overwriting of frame names by body names when saving an XML. Introduced in 3.1.4.
|
||||
13. Fixed bug in Python binding of :ref:`mj_saveModel`: ``buffer`` argument was documented as optional but was actually
|
||||
12. Defaults of lights were not being saved, now fixed.
|
||||
13. Prevent overwriting of frame names by body names when saving an XML. Introduced in 3.1.4.
|
||||
14. Fixed bug in Python binding of :ref:`mj_saveModel`: ``buffer`` argument was documented as optional but was actually
|
||||
not optional.
|
||||
|
||||
|
||||
|
||||
@@ -56,7 +56,10 @@ from mujoco.mjx._src.collision_primitive import plane_ellipsoid
|
||||
from mujoco.mjx._src.collision_primitive import plane_sphere
|
||||
from mujoco.mjx._src.collision_primitive import sphere_capsule
|
||||
from mujoco.mjx._src.collision_primitive import sphere_sphere
|
||||
from mujoco.mjx._src.collision_sdf import capsule_cylinder
|
||||
from mujoco.mjx._src.collision_sdf import capsule_ellipsoid
|
||||
from mujoco.mjx._src.collision_sdf import cylinder_cylinder
|
||||
from mujoco.mjx._src.collision_sdf import ellipsoid_cylinder
|
||||
from mujoco.mjx._src.collision_sdf import ellipsoid_ellipsoid
|
||||
from mujoco.mjx._src.collision_types import FunctionKey
|
||||
from mujoco.mjx._src.types import Contact
|
||||
@@ -82,8 +85,11 @@ _COLLISION_FUNC = {
|
||||
(GeomType.CAPSULE, GeomType.CAPSULE): capsule_capsule,
|
||||
(GeomType.CAPSULE, GeomType.BOX): capsule_convex,
|
||||
(GeomType.CAPSULE, GeomType.ELLIPSOID): capsule_ellipsoid,
|
||||
(GeomType.CAPSULE, GeomType.CYLINDER): capsule_cylinder,
|
||||
(GeomType.CAPSULE, GeomType.MESH): capsule_convex,
|
||||
(GeomType.ELLIPSOID, GeomType.ELLIPSOID): ellipsoid_ellipsoid,
|
||||
(GeomType.ELLIPSOID, GeomType.CYLINDER): ellipsoid_cylinder,
|
||||
(GeomType.CYLINDER, GeomType.CYLINDER): cylinder_cylinder,
|
||||
(GeomType.BOX, GeomType.BOX): convex_convex,
|
||||
(GeomType.BOX, GeomType.MESH): convex_convex,
|
||||
(GeomType.MESH, GeomType.MESH): convex_convex,
|
||||
|
||||
@@ -264,6 +264,29 @@ class EllipsoidCollisionTest(parameterized.TestCase):
|
||||
_assert_attr_eq(
|
||||
dx.contact, d.contact, field.name, 'ellipsoid-capsule', 1e-4)
|
||||
|
||||
_ELLIPSOID_CYLINDER = """
|
||||
<mujoco>
|
||||
<worldbody>
|
||||
<body>
|
||||
<geom size=".15 .05" type="cylinder"/>
|
||||
</body>
|
||||
<body pos="0 0 0.09">
|
||||
<freejoint/>
|
||||
<geom size=".15 .03 .05" type="ellipsoid"/>
|
||||
</body>
|
||||
</worldbody>
|
||||
</mujoco>
|
||||
"""
|
||||
|
||||
def test_ellipsoid_cylinder(self):
|
||||
"""Tests ellipsoid cylinder contact."""
|
||||
d, dx = _collide(self._ELLIPSOID_CYLINDER)
|
||||
d.contact.pos[0][2] = 0.04 # MJX finds the deepest point on the surface
|
||||
self.assertLess(dx.contact.dist[0], 0)
|
||||
for field in dataclasses.fields(Contact):
|
||||
_assert_attr_eq(
|
||||
dx.contact, d.contact, field.name, 'ellipsoid-cylinder', 1e-4)
|
||||
|
||||
|
||||
class CapsuleCollisionTest(parameterized.TestCase):
|
||||
_CAP_PLANE = """
|
||||
|
||||
@@ -85,6 +85,43 @@ def _ellipsoid(pos: jax.Array, size: jax.Array) -> jax.Array:
|
||||
return k0 * (k0 - 1.0) / (k1 + (k1 == 0.0) * 1e-12)
|
||||
|
||||
|
||||
@jax.custom_jvp
|
||||
def _cylinder(pos: jax.Array, size: jax.Array) -> jax.Array:
|
||||
a0 = jp.sqrt(pos[0] * pos[0] + pos[1] * pos[1]) - size[0]
|
||||
a1 = jp.abs(pos[2]) - size[1]
|
||||
b0 = jp.maximum(a0, 0)
|
||||
b1 = jp.maximum(a1, 0)
|
||||
return jp.minimum(jp.maximum(a0, a1), 0) + jp.sqrt(b0 * b0 + b1 * b1)
|
||||
|
||||
|
||||
def _cylinder_grad(x: jax.Array, size: jax.Array) -> jax.Array:
|
||||
"""Gradient of the cylinder SDF wrt query point and singularities removed."""
|
||||
c = jp.sqrt(x[0]*x[0]+x[1]*x[1])
|
||||
e = jp.abs(x[2])
|
||||
a = jp.array([c - size[0], e - size[1]])
|
||||
b = jp.array([jp.maximum(a[0], 0), jp.maximum(a[1], 0)])
|
||||
j = jp.argmax(a)
|
||||
bnorm = jp.sqrt(b[0]*b[0] + b[1]*b[1])
|
||||
bnorm += jp.allclose(bnorm, 0) * 1e-12
|
||||
grada = jp.array([
|
||||
x[0] / (c + jp.allclose(c, 0) * 1e-12),
|
||||
x[1] / (c + jp.allclose(c, 0) * 1e-12),
|
||||
x[2] / (e + jp.allclose(e, 0) * 1e-12),
|
||||
])
|
||||
gradm = jp.array([[grada[0], grada[1], 0], [0, 0, grada[2]]])
|
||||
gradb = grada * b[jp.array([0, 0, 1])] / bnorm
|
||||
return jp.where(a[j] < 0, gradm[j], gradb)
|
||||
|
||||
|
||||
@_cylinder.defjvp
|
||||
def cylinder_jvp(primals, tangents):
|
||||
x, y = primals
|
||||
x_dot, _ = tangents
|
||||
primal_out = _cylinder(x, y)
|
||||
tangent_out = jp.dot(_cylinder_grad(x, y), x_dot)
|
||||
return primal_out, tangent_out
|
||||
|
||||
|
||||
def _to_local(f: SDFFn, pos: jax.Array, mat: jax.Array)-> SDFFn:
|
||||
return lambda p: f(mat.T @ (p - pos))
|
||||
|
||||
@@ -134,29 +171,73 @@ def _gradient_descent(
|
||||
|
||||
|
||||
def _optim(
|
||||
d1, d2, info1: GeomInfo, info2: GeomInfo
|
||||
) -> Tuple[jax.Array, jax.Array, jax.Array]:
|
||||
d1, d2, info1: GeomInfo, info2: GeomInfo, x0: jax.Array,
|
||||
) -> Collision:
|
||||
"""Optimizes the clearance function."""
|
||||
d1 = functools.partial(d1, size=info1.size)
|
||||
d1 = _to_local(d1, info1.pos, info1.mat)
|
||||
d2 = functools.partial(d2, size=info2.size)
|
||||
d2 = _to_local(d2, info2.pos, info2.mat)
|
||||
fn = _clearance(d1, d2)
|
||||
_, pos = _gradient_descent(fn, 0.5 * (info1.pos + info2.pos), 10)
|
||||
_, pos = _gradient_descent(fn, x0, 10)
|
||||
dist = d1(pos) + d2(pos)
|
||||
n = jax.grad(d1)(pos)
|
||||
return pos, dist, n
|
||||
n = jax.grad(d1)(pos) - jax.grad(d2)(pos)
|
||||
return dist, pos, math.make_frame(n)
|
||||
|
||||
|
||||
@collider(ncon=1)
|
||||
def capsule_ellipsoid(c: GeomInfo, e: GeomInfo) -> Collision:
|
||||
""""Calculates contact between a capsule and an ellipsoid."""
|
||||
pos, dist, n = _optim(_capsule, _ellipsoid, c, e)
|
||||
return dist, pos, math.make_frame(n)
|
||||
x0 = 0.5 * (c.pos + e.pos)
|
||||
return _optim(_capsule, _ellipsoid, c, e, x0)
|
||||
|
||||
|
||||
@collider(ncon=2)
|
||||
def capsule_cylinder(ca: GeomInfo, cy: GeomInfo) -> Collision:
|
||||
""""Calculates contact between a capsule and a cylinder."""
|
||||
# TODO: improve robustness
|
||||
# Near sharp corners, the SDF might give the penetration depth with respect
|
||||
# to a surface that is not in collision. Possible solutions is to find the
|
||||
# contact points analytically or to change the SDF depending on the relative
|
||||
# pose of the bodies.
|
||||
mid = 0.5 * (ca.pos + cy.pos)
|
||||
vec = ca.mat[:, 2] * ca.size[1]
|
||||
x0 = jp.array([mid - vec, mid + vec])
|
||||
optim_ = functools.partial(_optim, _capsule, _cylinder, ca, cy)
|
||||
return jax.vmap(optim_)(x0)
|
||||
|
||||
|
||||
@collider(ncon=1)
|
||||
def ellipsoid_ellipsoid(e1: GeomInfo, e2: GeomInfo) -> Collision:
|
||||
""""Calculates contact between two ellipsoids."""
|
||||
pos, dist, n = _optim(_ellipsoid, _ellipsoid, e1, e2)
|
||||
return dist, pos, math.make_frame(n)
|
||||
x0 = 0.5 * (e1.pos + e2.pos)
|
||||
return _optim(_ellipsoid, _ellipsoid, e1, e2, x0)
|
||||
|
||||
|
||||
@collider(ncon=1)
|
||||
def ellipsoid_cylinder(e: GeomInfo, c: GeomInfo) -> Collision:
|
||||
""""Calculates contact between and ellipsoid and a cylinder."""
|
||||
x0 = 0.5 * (e.pos + c.pos)
|
||||
return _optim(_ellipsoid, _cylinder, e, c, x0)
|
||||
|
||||
|
||||
@collider(ncon=4)
|
||||
def cylinder_cylinder(c1: GeomInfo, c2: GeomInfo) -> Collision:
|
||||
""""Calculates contact between a cylinder and a cylinder."""
|
||||
# TODO: improve robustness
|
||||
# Near sharp corners, the SDF might give the penetration depth with respect
|
||||
# to a surface that is not in collision. Possible solutions is to find the
|
||||
# contact points analytically or to change the SDF depending on the relative
|
||||
# pose of the bodies.
|
||||
basis = math.make_frame(c2.pos - c1.pos)
|
||||
mid = 0.5 * (c1.pos + c2.pos)
|
||||
r = jp.maximum(c1.size[0], c2.size[0])
|
||||
x0 = jp.array([
|
||||
mid + r * basis[1],
|
||||
mid + r * basis[2],
|
||||
mid - r * basis[1],
|
||||
mid - r * basis[2],
|
||||
])
|
||||
optim_ = functools.partial(_optim, _cylinder, _cylinder, c1, c2)
|
||||
return jax.vmap(optim_)(x0)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user