diff --git a/doc/changelog.rst b/doc/changelog.rst index c1eb2f57..54230557 100644 --- a/doc/changelog.rst +++ b/doc/changelog.rst @@ -32,10 +32,11 @@ MJX 6. Fixed a bug in ``mjx.solve`` that was causing slow convergence when using ``mjSOL_NEWTON`` in :ref:`mjtSolver`. 7. Added support for :ref:`mjOption.impratio` to ``mjx.Model``. 8. Added support for cameras in ``mjx.Model`` and ``mjx.Data``. Fixes :github:issue:`1422`. +9. Added an implementation of broadphase using `top_k` and bounding spheres. Python bindings ^^^^^^^^^^^^^^^ -9. Fixed incorrect data types in the bindings for the ``geom``, ``vert``, ``elem``, and ``flex`` array members +10. Fixed incorrect data types in the bindings for the ``geom``, ``vert``, ``elem``, and ``flex`` array members of the ``mjContact`` struct, and all array members of the ``mjrContext`` struct. diff --git a/mjx/mujoco/mjx/_src/collision_base.py b/mjx/mujoco/mjx/_src/collision_base.py index c843dceb..fbad8609 100644 --- a/mjx/mujoco/mjx/_src/collision_base.py +++ b/mjx/mujoco/mjx/_src/collision_base.py @@ -44,6 +44,7 @@ CandidateSet = Dict[ class GeomInfo(PyTreeNode): """Collision info for a geom.""" + geom_id: jax.Array pos: jax.Array mat: jax.Array size: jax.Array diff --git a/mjx/mujoco/mjx/_src/collision_convex.py b/mjx/mujoco/mjx/_src/collision_convex.py index 21197c15..5428f53c 100644 --- a/mjx/mujoco/mjx/_src/collision_convex.py +++ b/mjx/mujoco/mjx/_src/collision_convex.py @@ -520,7 +520,7 @@ def plane_convex(plane: GeomInfo, convex: GeomInfo) -> Contact: def sphere_convex(sphere: GeomInfo, convex: GeomInfo) -> Contact: """Calculates contact between a sphere and a convex object.""" - faces = jp.take(convex.vert, convex.face, axis=0) + faces = convex.face normals = convex.facenorm # Put sphere in convex frame. @@ -583,7 +583,7 @@ def sphere_convex(sphere: GeomInfo, convex: GeomInfo) -> Contact: def capsule_convex(cap: GeomInfo, convex: GeomInfo) -> Contact: """Calculates contacts between a capsule and a convex object.""" # Get convex transformed normals, faces, and vertices. - faces = jp.take(convex.vert, convex.face, axis=0) + faces = convex.face normals = convex.facenorm # Put capsule in convex frame. @@ -669,12 +669,13 @@ def convex_convex(c1: GeomInfo, c2: GeomInfo) -> Contact: if c1.face is None or c2.face is None or c1.vert is None or c2.vert is None: raise AssertionError('Mesh info missing.') # pad face vertices so that we can broadcast between geom1 and geom2 - s1, s2 = c1.face.shape[-1], c2.face.shape[-1] - if s1 < s2: - face = jp.pad(c1.face, ((0, 0), (0, s2 - s1)), 'edge') + # face has shape (n_face, n_vert, 3) + nvert1, nvert2 = c1.face.shape[1], c2.face.shape[1] + if nvert1 < nvert2: + face = jp.pad(c1.face, ((0, 0), (0, nvert2 - nvert1), (0, 0)), 'edge') c1 = c1.replace(face=face) - elif s2 < s1: - face = jp.pad(c2.face, ((0, 0), (0, s1 - s2)), 'edge') + elif nvert2 < nvert1: + face = jp.pad(c2.face, ((0, 0), (0, nvert1 - nvert2), (0, 0)), 'edge') c2 = c2.replace(face=face) # ensure that the first object has fewer verts @@ -682,8 +683,8 @@ def convex_convex(c1: GeomInfo, c2: GeomInfo) -> Contact: if swapped: c1, c2 = c2, c1 - faces1 = jp.take(c1.vert, c1.face, axis=0) - faces2 = jp.take(c2.vert, c2.face, axis=0) + faces1 = c1.face + faces2 = c2.face to_local_pos = c2.mat.T @ (c1.pos - c2.pos) to_local_mat = c2.mat.T @ c1.mat diff --git a/mjx/mujoco/mjx/_src/collision_driver.py b/mjx/mujoco/mjx/_src/collision_driver.py index 721b4345..bdcf5275 100644 --- a/mjx/mujoco/mjx/_src/collision_driver.py +++ b/mjx/mujoco/mjx/_src/collision_driver.py @@ -20,6 +20,8 @@ import jax from jax import numpy as jp import mujoco from mujoco.mjx._src import collision_base +from mujoco.mjx._src import support +from mujoco.mjx._src import math # pylint: disable=g-importing-member from mujoco.mjx._src.collision_base import Candidate from mujoco.mjx._src.collision_base import CandidateSet @@ -181,11 +183,13 @@ def _pair_info( """Returns geom pair info for calculating collision.""" g1, g2 = jp.array(geom1), jp.array(geom2) info1 = GeomInfo( + g1, d.geom_xpos[g1], d.geom_xmat[g1], m.geom_size[g1], ) info2 = GeomInfo( + g2, d.geom_xpos[g2], d.geom_xmat[g2], m.geom_size[g2], @@ -236,6 +240,18 @@ def _body_pair_filter( return False +def _broadphase_enabled( + geom_types: Tuple[GeomType, GeomType], + n_pairs: int, + max_pairs: int, +) -> bool: + return ( + GeomType.PLANE not in geom_types + and max_pairs > -1 + and n_pairs > max_pairs + ) + + def _collide_geoms( m: Model, d: Data, @@ -264,18 +280,32 @@ def _collide_geoms( else: params.append(_dynamic_params(m, candidates)) - # call contact function + params = jax.tree_map(lambda *x: jp.concatenate(x), *params) g1, g2, in_axes = _pair_info(m, d, geom1, geom2) + + # Run a crude version of broadphase. + max_pairs = int(support.get_custom_numeric(m, 'max_geom_pairs')) + run_broadphase = _broadphase_enabled(geom_types, len(geom1), max_pairs) + n_pairs = max_pairs if run_broadphase else len(geom1) + if run_broadphase: + # broadphase over geom pairs, using bounding spheres + size1 = jp.max(m.geom_size[jp.array(geom1)], axis=-1) + size2 = jp.max(m.geom_size[jp.array(geom2)], axis=-1) + # TODO(btaba): consider re-using collision info for (sphere, sphere) + dists = jax.vmap(jp.linalg.norm)(g2.pos - g1.pos) - (size1 + size2) + _, idx = jax.lax.top_k(dists, k=n_pairs) + g1, g2, params = jax.tree_map( + lambda x, idx=idx: x[idx, ...], (g1, g2, params) + ) + + # call contact function res = jax.vmap(fn, in_axes=in_axes)(g1, g2) dist, pos, frame = jax.tree_map(jp.concatenate, res) - params = jax.tree_map(lambda *x: jp.concatenate(x), *params) - geom1, geom2 = jp.array(geom1), jp.array(geom2) # repeat params by the number of contacts per geom pair - n_repeat = dist.shape[-1] // geom1.shape[0] geom1, geom2, params = jax.tree_map( - lambda x: jp.repeat(x, n_repeat, axis=0), - (geom1, geom2, params), + lambda x: jp.repeat(x, fn.ncon, axis=0), # pytype: disable=attribute-error + (g1.geom_id, g2.geom_id, params), ) con = Contact( @@ -293,16 +323,6 @@ def _collide_geoms( return con -def _max_contact_points(m: Union[Model, mujoco.MjModel]) -> int: - """Returns the maximum number of contact points when set as a numeric.""" - for i in range(m.nnumeric): - name = m.names[m.name_numericadr[i] :].decode('utf-8').split('\x00', 1)[0] - if name == 'max_contact_points': - return int(m.numeric_data[m.numeric_adr[i]]) - - return -1 - - def collision_candidates(m: Union[Model, mujoco.MjModel]) -> CandidateSet: """Returns candidates for collision checking.""" candidate_set = {} @@ -351,14 +371,17 @@ def ncon(m: Union[Model, mujoco.MjModel]) -> int: return 0 candidates = collision_candidates(m) - max_count = _max_contact_points(m) + max_count = int(support.get_custom_numeric(m, 'max_contact_points')) + max_pairs = int(support.get_custom_numeric(m, 'max_geom_pairs')) count = 0 for k, v in candidates.items(): fn = get_collision_fn(k[0:2]) if fn is None: continue - count += len(v) * fn.ncon # pytype: disable=attribute-error + run_broadphase = _broadphase_enabled((k[0], k[1]), len(v), max_pairs) + n_pair = max_pairs if run_broadphase else len(v) + count += n_pair * fn.ncon # pytype: disable=attribute-error return min(max_count, count) if max_count > -1 else count @@ -380,7 +403,7 @@ def collision(m: Model, d: Data) -> Data: contact = jax.tree_map(lambda *x: jp.concatenate(x), *contacts) - max_contact_points = _max_contact_points(m) + max_contact_points = int(support.get_custom_numeric(m, 'max_contact_points')) if max_contact_points > -1 and contact.dist.shape[0] > max_contact_points: # get top-k contacts _, idx = jax.lax.top_k(-contact.dist, k=max_contact_points) diff --git a/mjx/mujoco/mjx/_src/collision_driver_test.py b/mjx/mujoco/mjx/_src/collision_driver_test.py index f0bca5e9..3b7aff1d 100644 --- a/mjx/mujoco/mjx/_src/collision_driver_test.py +++ b/mjx/mujoco/mjx/_src/collision_driver_test.py @@ -537,6 +537,53 @@ class TopKContactTest(absltest.TestCase): self.assertEqual(dx_all.contact.dist.shape, (3,)) self.assertEqual(dx_top_k.contact.dist.shape, (2,)) + _CAPSULES_MAX_PAIR = """ + + + + + + + + + + + + + + + + + + + + + + + + """ + + def test_max_pair(self): + """Tests contact culling before the collision functions were dispatched.""" + with jax.disable_jit(): + m = mujoco.MjModel.from_xml_string(self._CAPSULES_MAX_PAIR) + mx_top_k = mjx.put_model(m) + mx_all = mx_top_k.replace( + nnumeric=0, name_numericadr=np.array([]), numeric_data=np.array([]) + ) + d = mujoco.MjData(m) + dx = mjx.put_data(m, d) + + collision_jit_fn = jax.jit(mjx.collision) + kinematics_jit_fn = jax.jit(mjx.kinematics) + dx = kinematics_jit_fn(mx_all, dx) + + dx_all = collision_jit_fn(mx_all, dx) + dx_top_k = collision_jit_fn(mx_top_k, dx) + + self.assertEqual(dx_all.contact.dist.shape, (6,)) + self.assertEqual(dx_top_k.contact.dist.shape, (2,)) + if __name__ == '__main__': absltest.main() diff --git a/mjx/mujoco/mjx/_src/mesh.py b/mjx/mujoco/mjx/_src/mesh.py index bebca9ac..eb7a79c1 100644 --- a/mjx/mujoco/mjx/_src/mesh.py +++ b/mjx/mujoco/mjx/_src/mesh.py @@ -201,7 +201,8 @@ def _geom_mesh_kwargs( vert = np.array(tm_convex.vertices) face = _merge_coplanar(tm_convex) return { - 'geom_convex_face': face, + 'geom_convex_face': vert[face], + 'geom_convex_face_vert_idx': face, 'geom_convex_vert': vert, 'geom_convex_edge': _get_unique_edges(vert, face), 'geom_convex_facenormal': _get_face_norm(vert, face), diff --git a/mjx/mujoco/mjx/_src/mesh_test.py b/mjx/mujoco/mjx/_src/mesh_test.py index 7be8197a..0ff9785e 100644 --- a/mjx/mujoco/mjx/_src/mesh_test.py +++ b/mjx/mujoco/mjx/_src/mesh_test.py @@ -47,7 +47,7 @@ class GeomMeshKwargsTest(absltest.TestCase): # check face vertices map_ = {v: k for k, v in enumerate(vidx)} - h_face = np.vectorize(map_.get)(h['geom_convex_face']) + h_face = np.vectorize(map_.get)(h['geom_convex_face_vert_idx']) face_verts = sorted([tuple(sorted(set(s))) for s in h_face.tolist()]) expected_face_verts = sorted([ (0, 3, 4), (1, 3, 4), (0, 2, 4), (0, 1, 2, 3), (1, 2, 4)]) diff --git a/mjx/mujoco/mjx/_src/support.py b/mjx/mujoco/mjx/_src/support.py index 0bdc0d50..03966fa2 100644 --- a/mjx/mujoco/mjx/_src/support.py +++ b/mjx/mujoco/mjx/_src/support.py @@ -199,3 +199,13 @@ def local_to_global( pos = world_pos + math.rotate(local_pos, world_quat) mat = math.quat_to_mat(math.quat_mul(world_quat, local_quat)) return pos, mat + + +def get_custom_numeric(m: Union[Model, mujoco.MjModel], name: str) -> float: + """Returns a custom numeric given an MjModel or mjx.Model.""" + for i in range(m.nnumeric): + name_ = m.names[m.name_numericadr[i] :].decode('utf-8').split('\x00', 1)[0] + if name_ == name: + return m.numeric_data[m.numeric_adr[i]] + + return -1 diff --git a/mjx/mujoco/mjx/_src/support_test.py b/mjx/mujoco/mjx/_src/support_test.py index 253b6ec6..76606ca7 100644 --- a/mjx/mujoco/mjx/_src/support_test.py +++ b/mjx/mujoco/mjx/_src/support_test.py @@ -116,6 +116,20 @@ class SupportTest(parameterized.TestCase): np.testing.assert_almost_equal(qfrc, qfrc_expected, 6) + def test_custom_numeric(self): + xml = """ + + + + + + + """ + m = mujoco.MjModel.from_xml_string(xml) + self.assertEqual(support.get_custom_numeric(m, 'something'), -1) + self.assertEqual(support.get_custom_numeric(m, 'max_contact_points'), 15) + self.assertEqual(support.get_custom_numeric(m, 'max_geom_pairs'), 42) + if __name__ == '__main__': absltest.main() diff --git a/mjx/mujoco/mjx/test_data/shadow_hand/right_hand.xml b/mjx/mujoco/mjx/test_data/shadow_hand/right_hand.xml index b0795c28..a5907d8b 100644 --- a/mjx/mujoco/mjx/test_data/shadow_hand/right_hand.xml +++ b/mjx/mujoco/mjx/test_data/shadow_hand/right_hand.xml @@ -7,6 +7,7 @@ +