Add cylinder plane.
PiperOrigin-RevId: 625099227 Change-Id: I02cb86c3e8c7fd5914a082f1e8a3bae6274b3c8e
This commit is contained in:
committed by
Copybara-Service
parent
38fa5083eb
commit
0fa0437313
@@ -11,6 +11,12 @@ General
|
||||
attribute.
|
||||
|
||||
|
||||
MJX
|
||||
^^^
|
||||
|
||||
1. Add cylinder plane collisions.
|
||||
|
||||
|
||||
Version 3.1.4 (April 10th, 2024)
|
||||
--------------------------------
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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]
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user