From 0fa04373133eb45316b9e5c847939552d1039d44 Mon Sep 17 00:00:00 2001 From: Baruch Tabanpour Date: Mon, 15 Apr 2024 15:16:50 -0700 Subject: [PATCH] Add cylinder plane. PiperOrigin-RevId: 625099227 Change-Id: I02cb86c3e8c7fd5914a082f1e8a3bae6274b3c8e --- doc/changelog.rst | 6 ++ mjx/mujoco/mjx/_src/collision_driver.py | 2 + mjx/mujoco/mjx/_src/collision_driver_test.py | 44 ++++++++++++++ mjx/mujoco/mjx/_src/collision_primitive.py | 60 ++++++++++++++++++++ mjx/mujoco/mjx/_src/collision_sdf.py | 3 +- 5 files changed, 113 insertions(+), 2 deletions(-) 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]