Add cylinder plane.

PiperOrigin-RevId: 625099227
Change-Id: I02cb86c3e8c7fd5914a082f1e8a3bae6274b3c8e
This commit is contained in:
Baruch Tabanpour
2024-04-15 15:16:50 -07:00
committed by Copybara-Service
parent 38fa5083eb
commit 0fa0437313
5 changed files with 113 additions and 2 deletions
+6
View File
@@ -11,6 +11,12 @@ General
attribute.
MJX
^^^
1. Add cylinder plane collisions.
Version 3.1.4 (April 10th, 2024)
--------------------------------
+2
View File
@@ -33,6 +33,7 @@ from mujoco.mjx._src.collision_convex import plane_convex
from mujoco.mjx._src.collision_convex import sphere_convex
from mujoco.mjx._src.collision_primitive import capsule_capsule
from mujoco.mjx._src.collision_primitive import plane_capsule
from mujoco.mjx._src.collision_primitive import plane_cylinder
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
@@ -52,6 +53,7 @@ _COLLISION_FUNC = {
(GeomType.PLANE, GeomType.CAPSULE): plane_capsule,
(GeomType.PLANE, GeomType.BOX): plane_convex,
(GeomType.PLANE, GeomType.ELLIPSOID): plane_ellipsoid,
(GeomType.PLANE, GeomType.CYLINDER): plane_cylinder,
(GeomType.PLANE, GeomType.MESH): plane_convex,
(GeomType.SPHERE, GeomType.SPHERE): sphere_sphere,
(GeomType.SPHERE, GeomType.CAPSULE): sphere_capsule,
@@ -465,6 +465,50 @@ class CapsuleCollisionTest(parameterized.TestCase):
)
class CylinderTest(absltest.TestCase):
"""Tests the cylinder contact functions."""
_CYLINDER_PLANE = """
<mujoco>
<worldbody>
<geom size="40 40 40" type="plane"/>
<body pos="0 0 0.04">
<joint type="free"/>
<geom fromto="-0.1 0 0 0.1 0 0" size="0.05" type="cylinder"/>
</body>
</worldbody>
</mujoco>
"""
def test_cylinder_plane(self):
d, dx = _collide(self._CYLINDER_PLANE)
# cylinder is lying flat
np.testing.assert_array_less(dx.contact.dist[:2], 0)
np.testing.assert_array_less(-dx.contact.dist[2:], 0)
# sort position for comparison
idx = np.lexsort((dx.contact.pos[:, 0], dx.contact.pos[:, 1]))
dx = dx.tree_replace({'contact.pos': dx.contact.pos[idx]})
idx = np.lexsort((d.contact.pos[:, 0], d.contact.pos[:, 1]))
d.contact.pos[:] = d.contact.pos[idx]
# extract the contact points with penetration
c = jax.tree_map(lambda x: jp.take(x, jp.array([0, 1]), axis=0), dx.contact)
for field in dataclasses.fields(Contact):
_assert_attr_eq(c, d.contact, field.name, 'cylinder_plane', 1e-5)
# cylinder is vertical
xml = self._CYLINDER_PLANE.replace(
'<geom fromto="-0.1 0 0 0.1 0 0"', '<geom fromto="0 0 -0.1 0 0 0.1"')
xml = xml.replace('pos="0 0 0.04"', 'pos="0 0 0.095"')
d, dx = _collide(xml)
np.testing.assert_array_less(dx.contact.dist, 0)
for field in dataclasses.fields(Contact):
_assert_attr_eq(dx.contact, d.contact, field.name, 'cylinder_plane', 1e-5)
class ConvexTest(absltest.TestCase):
"""Tests the convex contact functions."""
@@ -78,6 +78,65 @@ def plane_ellipsoid(plane: GeomInfo, ellipsoid: GeomInfo) -> Contact:
)
def plane_cylinder(plane: GeomInfo, cylinder: GeomInfo) -> Contact:
"""Calculates one contact between an cylinder and a plane."""
n = plane.mat[:, 2]
axis = cylinder.mat[:, 2]
# make sure axis points towards plane
prjaxis = jp.dot(n, axis)
sign = -math.sign(prjaxis)
axis, prjaxis = axis * sign, prjaxis * sign
# compute normal distance to cylinder center
dist0 = jp.dot(cylinder.pos - plane.pos, n)
# remove component of -normal along axis, compute length
vec = axis * prjaxis - n
len_ = math.norm(vec)
vec = jp.where(
len_ < 1e-12,
# disk parallel to plane: pick x-axis of cylinder, scale by radius
cylinder.mat[:, 0] * cylinder.size[0],
# general configuration: normalize vector, scale by radius
vec / len_ * cylinder.size[0]
)
# project vector on normal
prjvec = jp.dot(vec, n)
# scale axis by half-length
axis *= cylinder.size[1]
prjaxis *= cylinder.size[1]
# compute sideways vector: vec1
prjvec1 = -prjvec * 0.5
vec1 = math.normalize(jp.cross(vec, axis)) * cylinder.size[0]
vec1 *= jp.sqrt(3.0) * 0.5
# disk parallel to plane
d1 = dist0 + prjaxis + prjvec
d2 = dist0 + prjaxis + prjvec1
dist = jp.array([d1, d2, d2])
pos = cylinder.pos + axis + jp.array([
vec - n * d1 * 0.5,
vec1 + vec * -0.5 - n * d2 * 0.5,
-vec1 + vec * -0.5 - n * d2 * 0.5,
])
# cylinder parallel to plane
cond = jp.abs(prjaxis) < 1e-3
d3 = dist0 - prjaxis + prjvec
dist = jp.where(cond, dist.at[1].set(d3), dist)
pos = jp.where(
cond, pos.at[1].set(cylinder.pos + vec - axis - n * d3 * 0.5), pos
)
frame = jp.stack([math.make_frame(n)] * 3, axis=0)
return dist, pos, frame
def _sphere_sphere(
pos1: jax.Array, radius1: jax.Array, pos2: jax.Array, radius2: jax.Array
) -> Contact:
@@ -135,6 +194,7 @@ def capsule_capsule(cap1: GeomInfo, cap2: GeomInfo) -> Contact:
plane_sphere.ncon = 1
plane_capsule.ncon = 2
plane_ellipsoid.ncon = 1
plane_cylinder.ncon = 3
sphere_sphere.ncon = 1
sphere_capsule.ncon = 1
capsule_capsule.ncon = 1
+1 -2
View File
@@ -33,8 +33,7 @@ from mujoco.mjx._src.collision_base import GeomInfo
from mujoco.mjx._src.dataclasses import PyTreeNode
# pylint: enable=g-importing-member
# the objective function, inputs: pos (that we optimize for) and size (known)
# the SDF function takes position in, and returns a distance or objective
SDFFn = Callable[[jax.Array], jax.Array]