MJX (impl='warp'): nconmax -> naconmax

PiperOrigin-RevId: 915407157
Change-Id: Ie43c06164d6d942c2a4bb60b3d09fcb5c4c87a35
This commit is contained in:
Taylor Howell
2026-05-14 06:51:53 -07:00
committed by Copybara-Service
parent 9c372fbbea
commit 072125c49c
5 changed files with 12 additions and 33 deletions
+3
View File
@@ -59,6 +59,9 @@ General
Negative ``margin`` values are now permitted (corresponding to ``gap > margin`` under the old semantics). The
constraint ``margin + gap >= 0`` should be maintained to ensure valid collision detection.
- MJX: Removed the deprecated ``nconmax`` argument from ``mjx.make_data`` and ``mjx.put_data`` in favor of
``naconmax``.
Version 3.8.1 (May 11, 2026)
----------------------------
+2 -26
View File
@@ -828,7 +828,6 @@ def make_data(
device: Optional[jax.Device] = None,
impl: Optional[Union[str, types.Impl]] = None,
_full_compat: bool = False, # pylint: disable=invalid-name
nconmax: Optional[int] = None,
naconmax: Optional[int] = None,
naccdmax: Optional[int] = None,
njmax: Optional[int] = None,
@@ -840,15 +839,10 @@ def make_data(
m: the model to use
device: which device to use - if unspecified picks the default device
impl: implementation to use ('jax', 'warp')
nconmax: maximum number of contacts to allocate for warp across all worlds
Since the number of worlds is **not** pre-defined in JAX, we use the
`nconmax` argument to set the upper bound for the number of contacts
across all worlds. In MuJoCo Warp, the analgous field is called
`naconmax`.
naconmax: maximum number of contacts to allocate for warp across all worlds
Since the number of worlds is **not** pre-defined in JAX, we use the
`naconmax` argument to set the upper bound for the number of contacts
across all worlds, rather than the `nconmax` argument from MuJoCo Warp.
across all worlds.
naccdmax: maximum number of contacts for GJK collision detection across all
worlds. Since the number of worlds is **not** pre-defined in JAX, we use the
`naccdmax` argument to set the upper bound for the number of contacts
@@ -864,14 +858,7 @@ def make_data(
Raises:
ValueError: if the model's impl does not match the make_data impl
NotImplementedError: if the impl is not implemented yet
DeprecationWarning: if nconmax is used
"""
if nconmax is not None:
warnings.warn(
'nconmax will be deprecated in mujoco-mjx>=3.5. Use naconmax instead.',
DeprecationWarning,
stacklevel=2,
)
impl, device = _resolve_impl_and_device(impl, device)
@@ -886,7 +873,6 @@ def make_data(
return _make_data_cpp(m, device, keepalive_refs=keepalive_refs)
elif impl == types.Impl.WARP:
_check_warp_installed()
naconmax = nconmax if naconmax is None else naconmax
return _make_data_warp(m, device, naconmax, naccdmax, njmax)
raise NotImplementedError(
@@ -1187,7 +1173,6 @@ def put_data(
d: mujoco.MjData,
device: Optional[jax.Device] = None,
impl: Optional[Union[str, types.Impl]] = None,
nconmax: Optional[int] = None,
naconmax: Optional[int] = None,
njmax: Optional[int] = None,
dummy_arg_for_batching: Optional[jax.Array] = None,
@@ -1200,11 +1185,10 @@ def put_data(
d: the data to put on device
device: which device to use - if unspecified picks the default device
impl: implementation to use ('jax', 'warp')
nconmax: maximum number of contacts to allocate for warp
naconmax: maximum number of contacts to allocate for warp across all worlds
Since the number of worlds is **not** pre-defined in JAX, we use the
`naconmax` argument to set the upper bound for the number of contacts
across all worlds, rather than the `nconmax` argument from MuJoCo Warp.
across all worlds.
njmax: maximum number of constraints to allocate for warp
dummy_arg_for_batching: dummy argument to use for batching in cpp
implementation
@@ -1213,14 +1197,7 @@ def put_data(
Returns:
an mjx.Data placed on device
DeprecationWarning: if nconmax is used
"""
if nconmax is not None:
warnings.warn(
'nconmax will be deprecated in mujoco-mjx>=3.5. Use naconmax instead.',
DeprecationWarning,
stacklevel=2,
)
impl, device = _resolve_impl_and_device(impl, device)
if impl == types.Impl.JAX:
@@ -1235,7 +1212,6 @@ def put_data(
)
elif impl == types.Impl.WARP:
_check_warp_installed()
naconmax = nconmax if naconmax is None else naconmax
return _put_data_warp(m, d, device, naconmax, njmax)
raise NotImplementedError(
+1 -1
View File
@@ -475,7 +475,7 @@ class DataIOTest(parameterized.TestCase):
if not mjx_io.has_cuda_gpu_device():
self.skipTest('No CUDA GPU device.')
m = mujoco.MjModel.from_xml_string(_MULTIPLE_CONVEX_OBJECTS)
d = mjx.make_data(m, impl='warp', nconmax=9, njmax=23)
d = mjx.make_data(m, impl='warp', naconmax=9, njmax=23)
self.assertEqual(d._impl.contact__dist.shape[0], 9)
self.assertEqual(d._impl.efc__pos.shape[0], 23)
+4 -4
View File
@@ -329,7 +329,7 @@ def _check_leading_dim(
path: jax.tree_util.KeyPath,
leaf: Any,
expected_batch_dim: int,
expected_nconmax: int,
expected_naconmax: int,
expected_njmax: int,
):
"""Asserts that the batch dimension of a leaf node matches the expected batch dimension."""
@@ -345,11 +345,11 @@ def _check_leading_dim(
if (
not has_batch_dim
and attr.startswith('contact__')
and leaf.shape[0] != expected_nconmax
and leaf.shape[0] != expected_naconmax
):
raise ValueError(
f'Leaf node leading dim ({leaf.shape[0]}) does not match nconmax'
f' ({expected_nconmax}) for field {attr}.'
f'Leaf node leading dim ({leaf.shape[0]}) does not match naconmax'
f' ({expected_naconmax}) for field {attr}.'
)
if (
not has_batch_dim
+2 -2
View File
@@ -41,10 +41,10 @@ def assert_attr_eq(a, b, attr):
def make_data(
m: mujoco.MjModel, worldid: int, nconmax: int = 1_000, njmax: int = 200
m: mujoco.MjModel, worldid: int, naconmax: int = 1_000, njmax: int = 200
):
"""Make data for a given worldid using keyframes when available."""
dx = mjx.make_data(m, impl='warp', nconmax=nconmax, njmax=njmax)
dx = mjx.make_data(m, impl='warp', naconmax=naconmax, njmax=njmax)
rng = jax.random.PRNGKey(worldid)
rng, key = jax.random.split(rng)