diff --git a/mjx/mujoco/mjx/_src/collision_driver.py b/mjx/mujoco/mjx/_src/collision_driver.py index b5f58130..dde3fe92 100644 --- a/mjx/mujoco/mjx/_src/collision_driver.py +++ b/mjx/mujoco/mjx/_src/collision_driver.py @@ -336,16 +336,28 @@ def collision_candidates(m: Union[Model, mujoco.MjModel]) -> CandidateSet: def ncon(m: Model) -> int: """Returns the number of contacts computed in MJX given a model.""" + if m.opt.disableflags & DisableBit.CONTACT: + return 0 + candidates = collision_candidates(m) max_count = _max_contact_points(m) - count = sum([ - len(v) * get_collision_fn(k[0:2]).ncon for k, v in candidates.items() # pytype: disable=attribute-error - ]) + + 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 + return min(max_count, count) if max_count > -1 else count def collision(m: Model, d: Data) -> Data: """Collides geometries.""" + ncon_ = ncon(m) + if ncon_ == 0: + return d.replace(contact=Contact.zero(), ncon=0) + candidate_set = collision_candidates(m) contacts = [] @@ -354,7 +366,7 @@ def collision(m: Model, d: Data) -> Data: contacts.append(_collide_geoms(m, d, geom_types, candidates)) if not contacts: - return d.replace(contact=Contact.zero(), ncon=0) + raise RuntimeError('No contacts found.') contact = jax.tree_map(lambda *x: jp.concatenate(x), *contacts) @@ -364,10 +376,13 @@ def collision(m: Model, d: Data) -> Data: _, idx = jax.lax.top_k(-contact.dist, k=max_contact_points) contact = jax.tree_map(lambda x, idx=idx: jp.take(x, idx, axis=0), contact) - ncon_ = contact.dist.shape[0] + if ncon_ != contact.dist.shape[0]: + raise RuntimeError('Number of contacts does not match ncon.') + + # TODO(robotics-simulation): move this logic to device_put ns = d.ne + d.nf + d.nl + contact = contact.replace(efc_address=np.arange(ns, ns + ncon_ * 4, 4)) # TODO(robotics-simulation): add support for other friction dimensions - contact = contact.replace(efc_address=np.arange(ns, ns + d.ncon * 4, 4)) contact = contact.replace(dim=3 * np.ones(ncon_, dtype=np.int32)) return d.replace(contact=contact, ncon=ncon_) diff --git a/mjx/mujoco/mjx/_src/collision_driver_test.py b/mjx/mujoco/mjx/_src/collision_driver_test.py index e31b95c8..b735e014 100644 --- a/mjx/mujoco/mjx/_src/collision_driver_test.py +++ b/mjx/mujoco/mjx/_src/collision_driver_test.py @@ -24,9 +24,12 @@ import jax import jax.numpy as jp import mujoco from mujoco import mjx +from mujoco.mjx._src import collision_driver +from mujoco.mjx._src import test_util # pylint: disable=g-importing-member from mujoco.mjx._src.types import Contact from mujoco.mjx._src.types import Data +from mujoco.mjx._src.types import DisableBit from mujoco.mjx._src.types import Model # pylint: enable=g-importing-member import numpy as np @@ -450,6 +453,29 @@ class BodyPairFilterTest(absltest.TestCase): self.assertEqual(dx.contact.pos.shape[0], 1) +class NconTest(parameterized.TestCase): + """Tests ncon.""" + + def test_ncon(self): + m = test_util.load_test_file('ant.xml') + d = mujoco.MjData(m) + d.qpos[2] = 0.0 + + mx = mjx.device_put(m) + ncon = collision_driver.ncon(mx) + self.assertEqual(ncon, 4) + + def test_disable_contact(self): + m = test_util.load_test_file('ant.xml') + d = mujoco.MjData(m) + d.qpos[2] = 0.0 + + m.opt.disableflags = m.opt.disableflags | DisableBit.CONTACT + mx = mjx.device_put(m) + ncon = collision_driver.ncon(mx) + self.assertEqual(ncon, 0) + + class TopKContactTest(absltest.TestCase): """Tests top-k contacts."""