Add naconmax, deprecation warning for nconmax.

PiperOrigin-RevId: 825018530
Change-Id: I344b296ac4353592d65a85b826850eaa7df5d39a
This commit is contained in:
Baruch Tabanpour
2025-10-28 06:58:14 -07:00
committed by Copybara-Service
parent 2bdea1a0aa
commit d001bd6525
4 changed files with 43 additions and 13 deletions
+3
View File
@@ -39,6 +39,9 @@ General
**Migration:** Replace ``meshdir`` and ``texturedir`` with ``compiler.meshdir`` and ``compiler.texturedir``.
2. Remove ``_full_compat`` from ``mjx.put_data`` and ``mjx.put_model``.
3. ``nconmax`` and ``njmax`` fields in ``mjx.make_data`` now default to ``None`` instead of -1. ``nconmax`` will be deprecated
in favor of ``naconmax`` in a future release.
3. Joint decorators and spatial tendons which have limits defined and whose current value (angle or length) exceeds the
limit, are recolored by using the :ref:`constraint impedance<soParameters>` :math:`d` to mix the existing color with
+34 -7
View File
@@ -813,7 +813,7 @@ def _get_nested_attr(obj: Any, attr_name: str, split: str) -> Any:
def _make_data_warp(
m: Union[types.Model, mujoco.MjModel],
device: Optional[jax.Device] = None,
nconmax: Optional[int] = None,
naconmax: Optional[int] = None,
njmax: Optional[int] = None,
) -> types.Data:
"""Allocate and initialize Data for the Warp implementation."""
@@ -827,7 +827,7 @@ def _make_data_warp(
raise RuntimeError('Warp is not installed.')
with wp.ScopedDevice('cpu'): # pylint: disable=undefined-variable
dw = mjwp.make_data(m, nworld=1, naconmax=nconmax, njmax=njmax) # pylint: disable=undefined-variable
dw = mjwp.make_data(m, nworld=1, naconmax=naconmax, njmax=njmax) # pylint: disable=undefined-variable
fields = _make_data_public_fields(m)
for k in fields:
@@ -862,7 +862,7 @@ def _make_data_warp(
# TODO(robotics-simulation): remove this warmup compilation once warp
# stops unloading modules during XLA graph capture for tile kernels.
# pylint: disable=undefined-variable
dw = mjwp.make_data(m, nworld=1, naconmax=nconmax, njmax=njmax)
dw = mjwp.make_data(m, nworld=1, naconmax=naconmax, njmax=njmax)
mw = mjwp.put_model(m)
_ = mjwp.step(mw, dw)
# pylint: enable=undefined-variable
@@ -877,6 +877,7 @@ def make_data(
impl: Optional[Union[str, types.Impl]] = None,
_full_compat: bool = False, # pylint: disable=invalid-name
nconmax: Optional[int] = None,
naconmax: Optional[int] = None,
njmax: Optional[int] = None,
) -> types.Data:
"""Allocate and initialize Data.
@@ -890,6 +891,10 @@ def make_data(
`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.
njmax: maximum number of constraints to allocate for warp across all worlds
Returns:
@@ -898,7 +903,15 @@ 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)
if isinstance(m, types.Model) and m.impl != impl:
@@ -912,7 +925,8 @@ def make_data(
elif impl == types.Impl.C:
return _make_data_c(m, device)
elif impl == types.Impl.WARP:
return _make_data_warp(m, device, nconmax, njmax)
naconmax = nconmax if naconmax is None else naconmax
return _make_data_warp(m, device, naconmax, njmax)
raise NotImplementedError(
f'make_data for implementation "{impl}" not implemented yet.'
@@ -1202,8 +1216,9 @@ def put_data(
d: mujoco.MjData,
device: Optional[jax.Device] = None,
impl: Optional[Union[str, types.Impl]] = None,
nconmax: int = -1,
njmax: int = -1,
nconmax: Optional[int] = None,
naconmax: Optional[int] = None,
njmax: Optional[int] = None,
) -> types.Data:
"""Puts mujoco.MjData onto a device, resulting in mjx.Data.
@@ -1213,12 +1228,24 @@ def put_data(
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.
njmax: maximum number of constraints to allocate for warp
Returns:
an mjx.Data placed on device
DeprecationWarning: if nconmax is used
"""
del nconmax, njmax
del njmax
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:
return _put_data_jax(m, d, device)
+4 -4
View File
@@ -42,8 +42,8 @@ _WP_KERNEL_CACHE_DIR = flags.DEFINE_string(
None,
'Path to the Warp kernel cache directory.',
)
_NCONMAX = flags.DEFINE_integer(
'nconmax',
_NACONMAX = flags.DEFINE_integer(
'naconmax',
None,
'Maximum number of contacts to simulate, warp only.',
)
@@ -91,11 +91,11 @@ def _main(argv: Sequence[str]) -> None:
if _IMPL.value == 'warp':
# TODO(btaba): use put_data.
dx = mjx.make_data(
m, impl=_IMPL.value, nconmax=_NCONMAX.value, njmax=_NJMAX.value
m, impl=_IMPL.value, naconmax=_NACONMAX.value, njmax=_NJMAX.value
)
else:
dx = mjx.put_data(
m, d, impl=_IMPL.value, nconmax=_NCONMAX.value, njmax=_NJMAX.value
m, d, impl=_IMPL.value, naconmax=_NACONMAX.value, njmax=_NJMAX.value
)
print(f'Default backend: {jax.default_backend()}')
+2 -2
View File
@@ -43,8 +43,8 @@ _FUNCTION = flags.DEFINE_string(
_NSTEP = flags.DEFINE_integer('nstep', 1000, 'number of steps per rollout')
_NENV = flags.DEFINE_integer('nenv', 8192, 'number of environments to simulate')
_UNROLL = flags.DEFINE_integer('unroll', 4, 'number of steps to unroll')
_NCONMAX = flags.DEFINE_integer('nconmax', 30_000, 'max contacts')
_NJMAX = flags.DEFINE_integer('njmax', 10, 'max constraints per world')
_NACONMAX = flags.DEFINE_integer('naconmax', 30_000, 'max contacts')
_NJMAX = flags.DEFINE_integer('njmax', 10, 'max constraints')
_WP_KERNEL_CACHE_DIR = flags.DEFINE_string(
'wp_kernel_cache_dir',
None,