From 52478e15c34ba764c70d2037e584e8ef0f835ece Mon Sep 17 00:00:00 2001 From: Tom Power Date: Wed, 12 Nov 2025 02:08:18 -0800 Subject: [PATCH] Fix issue where contact pairs supported by C back-end are unsupported by MJX:C PiperOrigin-RevId: 831288268 Change-Id: Ied69be2d69a4cf0c8190bfaf8b6428e1dd8c8a70 --- mjx/mujoco/mjx/_src/collision_driver.py | 23 ++++++++++++++++++--- mjx/mujoco/mjx/_src/io.py | 8 ++++---- mjx/mujoco/mjx/_src/io_test.py | 27 +++++++++++++++++++++++++ 3 files changed, 51 insertions(+), 7 deletions(-) diff --git a/mjx/mujoco/mjx/_src/collision_driver.py b/mjx/mujoco/mjx/_src/collision_driver.py index 158be7a2..0f15524b 100644 --- a/mjx/mujoco/mjx/_src/collision_driver.py +++ b/mjx/mujoco/mjx/_src/collision_driver.py @@ -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) diff --git a/mjx/mujoco/mjx/_src/io.py b/mjx/mujoco/mjx/_src/io.py index 22ab4e24..cec570a1 100644 --- a/mjx/mujoco/mjx/_src/io.py +++ b/mjx/mujoco/mjx/_src/io.py @@ -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) diff --git a/mjx/mujoco/mjx/_src/io_test.py b/mjx/mujoco/mjx/_src/io_test.py index a52867ba..c9eb7cb8 100644 --- a/mjx/mujoco/mjx/_src/io_test.py +++ b/mjx/mujoco/mjx/_src/io_test.py @@ -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(""" + + + + + + + + + + + + + + + + """) + + 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."""