Fix issue where contact pairs supported by C back-end are unsupported by MJX:C
PiperOrigin-RevId: 831288268 Change-Id: Ied69be2d69a4cf0c8190bfaf8b6428e1dd8c8a70
This commit is contained in:
committed by
Copybara-Service
parent
414d6d80ac
commit
52478e15c3
@@ -74,12 +74,14 @@ from mujoco.mjx._src.types import Data
|
||||
from mujoco.mjx._src.types import DataJAX
|
||||
from mujoco.mjx._src.types import DisableBit
|
||||
from mujoco.mjx._src.types import GeomType
|
||||
from mujoco.mjx._src.types import Impl
|
||||
from mujoco.mjx._src.types import Model
|
||||
from mujoco.mjx._src.types import ModelJAX
|
||||
from mujoco.mjx._src.types import OptionJAX
|
||||
# pylint: enable=g-importing-member
|
||||
import numpy as np
|
||||
|
||||
|
||||
# pair-wise collision functions
|
||||
_COLLISION_FUNC = {
|
||||
(GeomType.PLANE, GeomType.SPHERE): plane_sphere,
|
||||
@@ -111,6 +113,8 @@ _COLLISION_FUNC = {
|
||||
(GeomType.MESH, GeomType.MESH): convex_convex,
|
||||
}
|
||||
|
||||
# Maximum constraint dimension for collision functions.
|
||||
_MAX_NCON = 8
|
||||
|
||||
# geoms for which we ignore broadphase
|
||||
_GEOM_NO_BROADPHASE = {GeomType.HFIELD, GeomType.PLANE}
|
||||
@@ -341,8 +345,13 @@ def _numeric(m: Union[Model, mujoco.MjModel], name: str) -> int:
|
||||
return int(m.numeric_data[id_]) if id_ >= 0 else -1
|
||||
|
||||
|
||||
def make_condim(m: Union[Model, mujoco.MjModel]) -> np.ndarray:
|
||||
def make_condim(
|
||||
m: Union[Model, mujoco.MjModel], impl: Impl = Impl.JAX
|
||||
) -> np.ndarray:
|
||||
"""Returns the dims of the contacts for a Model."""
|
||||
if impl not in (Impl.JAX, Impl.C):
|
||||
raise ValueError('make_condim only supports JAX and C backends.')
|
||||
|
||||
if isinstance(m, mujoco.MjModel):
|
||||
sdf_initpoints = m.opt.sdf_initpoints
|
||||
elif isinstance(m.opt._impl, OptionJAX):
|
||||
@@ -377,8 +386,16 @@ def make_condim(m: Union[Model, mujoco.MjModel]) -> np.ndarray:
|
||||
if k.types[1] == mujoco.mjtGeom.mjGEOM_SDF:
|
||||
ncon = sdf_initpoints
|
||||
else:
|
||||
func = _COLLISION_FUNC[k.types]
|
||||
ncon = func.ncon # pytype: disable=attribute-error
|
||||
func = _COLLISION_FUNC.get(k.types, None)
|
||||
if func is not None:
|
||||
ncon = func.ncon # pytype: disable=attribute-error
|
||||
elif impl == Impl.C:
|
||||
ncon = _MAX_NCON
|
||||
else:
|
||||
raise ValueError(
|
||||
f'Collision function not found for geom types {k.types[0]},',
|
||||
f'{k.types[1]}'
|
||||
)
|
||||
num_contacts = condim_counts.get(k.condim, 0) + ncon * v
|
||||
if max_contact_points > -1:
|
||||
num_contacts = min(max_contact_points, num_contacts)
|
||||
|
||||
@@ -602,7 +602,7 @@ def _make_data_jax(
|
||||
device: Optional[jax.Device] = None,
|
||||
) -> types.Data:
|
||||
"""Allocate and initialize Data for the JAX implementation."""
|
||||
dim = collision_driver.make_condim(m)
|
||||
dim = collision_driver.make_condim(m, impl=types.Impl.JAX)
|
||||
efc_type = constraint.make_efc_type(m, dim)
|
||||
ne, nf, nl, nc = constraint.counts(efc_type)
|
||||
ncon, nefc = dim.size, ne + nf + nl + nc
|
||||
@@ -692,7 +692,7 @@ def _make_data_c(
|
||||
# TODO(stunya): The C implementation should not use static dimensions, and
|
||||
# the backend implementation details should be kept hidden from JAX
|
||||
# altogether.
|
||||
dim = collision_driver.make_condim(m)
|
||||
dim = collision_driver.make_condim(m, impl=types.Impl.C)
|
||||
efc_type = constraint.make_efc_type(m, dim)
|
||||
efc_address = constraint.make_efc_address(m, dim, efc_type)
|
||||
ne, nf, nl, nc = constraint.counts(efc_type)
|
||||
@@ -994,7 +994,7 @@ def _put_data_jax(
|
||||
m: mujoco.MjModel, d: mujoco.MjData, device: Optional[jax.Device] = None
|
||||
) -> types.Data:
|
||||
"""Puts mujoco.MjData onto a device, resulting in mjx.Data."""
|
||||
dim = collision_driver.make_condim(m)
|
||||
dim = collision_driver.make_condim(m, impl=types.Impl.JAX)
|
||||
efc_type = constraint.make_efc_type(m, dim)
|
||||
efc_address = constraint.make_efc_address(m, dim, efc_type)
|
||||
ne, nf, nl, nc = constraint.counts(efc_type)
|
||||
@@ -1121,7 +1121,7 @@ def _put_data_c(
|
||||
"""Puts mujoco.MjData onto a device, resulting in mjx.Data."""
|
||||
# TODO(stunya): ncon, nefc should potentially be jax.Array, and contact/efc
|
||||
# should not be materialized in JAX.
|
||||
dim = collision_driver.make_condim(m)
|
||||
dim = collision_driver.make_condim(m, impl=types.Impl.C)
|
||||
efc_type = constraint.make_efc_type(m, dim)
|
||||
efc_address = constraint.make_efc_address(m, dim, efc_type)
|
||||
ne, nf, nl, nc = constraint.counts(efc_type)
|
||||
|
||||
@@ -331,6 +331,33 @@ class ModelIOTest(parameterized.TestCase):
|
||||
|
||||
_ = jax.tree.map_with_path(check_ndim, mx)
|
||||
|
||||
@parameterized.parameters('c', 'jax')
|
||||
def test_unsupported_contact_types(self, impl):
|
||||
"""Tests that unsupported contact types raise an error."""
|
||||
m = mujoco.MjModel.from_xml_string("""
|
||||
<mujoco>
|
||||
<asset>
|
||||
<mesh name="box" vertex="-1 -1 -1 1 -1 -1 1 1 -1 1 1 1 1 -1 1 -1 1 -1 -1 1 1 -1 -1 1" scale="1 1 .1"/>
|
||||
</asset>
|
||||
<worldbody>
|
||||
<body name="meshbox">
|
||||
<freejoint/>
|
||||
<geom type="mesh" mesh="box" pos="0 0 -.15" euler="3 7 30"/>
|
||||
</body>
|
||||
<body name="cylinder">
|
||||
<freejoint/>
|
||||
<geom type="cylinder" size="1.0 0.1"/>
|
||||
</body>
|
||||
</worldbody>
|
||||
</mujoco>
|
||||
""")
|
||||
|
||||
if impl == 'jax':
|
||||
with self.assertRaises(ValueError):
|
||||
mjx.make_data(m, impl=impl)
|
||||
if impl == 'c':
|
||||
mjx.make_data(m, impl=impl)
|
||||
|
||||
|
||||
class DataIOTest(parameterized.TestCase):
|
||||
"""IO tests for mjx.Data."""
|
||||
|
||||
Reference in New Issue
Block a user