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:
Tom Power
2025-11-12 02:08:18 -08:00
committed by Copybara-Service
parent 414d6d80ac
commit 52478e15c3
3 changed files with 51 additions and 7 deletions
+20 -3
View File
@@ -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)
+4 -4
View File
@@ -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)
+27
View File
@@ -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."""