diff --git a/doc/changelog.rst b/doc/changelog.rst
index ab454f4c..a964c9fb 100644
--- a/doc/changelog.rst
+++ b/doc/changelog.rst
@@ -11,6 +11,12 @@ General
attribute.
+MJX
+^^^
+
+1. Add cylinder plane collisions.
+
+
Version 3.1.4 (April 10th, 2024)
--------------------------------
diff --git a/mjx/mujoco/mjx/_src/collision_driver.py b/mjx/mujoco/mjx/_src/collision_driver.py
index f4da34ef..7fd905b4 100644
--- a/mjx/mujoco/mjx/_src/collision_driver.py
+++ b/mjx/mujoco/mjx/_src/collision_driver.py
@@ -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,
diff --git a/mjx/mujoco/mjx/_src/collision_driver_test.py b/mjx/mujoco/mjx/_src/collision_driver_test.py
index d070ed7b..d501be5d 100644
--- a/mjx/mujoco/mjx/_src/collision_driver_test.py
+++ b/mjx/mujoco/mjx/_src/collision_driver_test.py
@@ -465,6 +465,50 @@ class CapsuleCollisionTest(parameterized.TestCase):
)
+class CylinderTest(absltest.TestCase):
+ """Tests the cylinder contact functions."""
+
+ _CYLINDER_PLANE = """
+
+
+
+
+
+
+
+
+
+ """
+
+ 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(
+ ' 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
diff --git a/mjx/mujoco/mjx/_src/collision_sdf.py b/mjx/mujoco/mjx/_src/collision_sdf.py
index 914a8e5b..999b3e8a 100644
--- a/mjx/mujoco/mjx/_src/collision_sdf.py
+++ b/mjx/mujoco/mjx/_src/collision_sdf.py
@@ -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]