Add contact disable flag to collision driver.
PiperOrigin-RevId: 575245746 Change-Id: If2ea8a5de6f6c751729c9b8dd5efc42e7bdd51ea
This commit is contained in:
committed by
Copybara-Service
parent
01a932e052
commit
ebe081f27e
@@ -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_)
|
||||
|
||||
@@ -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."""
|
||||
|
||||
|
||||
Reference in New Issue
Block a user