Previously mjx dataclasses would return np.ndarray as bytes for jax tracing to hash since they are not inherently hashable. This meant that any tracing that was cached would hold on to a copy of the numpy arrays in the dataclass. Children such as mjx.Model that store large numpy arrays would end up duplicating that data 3+ times in some cases. This eliminates O(N * array_size) memory duplication across N cached pytree traces, saving a lot of memory on models with heavy mesh/texture data. PiperOrigin-RevId: 880985898 Change-Id: I58d4e91fda4112e4c633818f92907d20138e962a
MuJoCo XLA (MJX)
This package is a re-implementation of the MuJoCo physics engine in JAX. This library is developed and maintained by Google DeepMind, and is kept up-to-date with the latest developments in MuJoCo itself.
The mujoco-mjx package is API-compatible with MuJoCo, but is missing some
features found in MuJoCo. See our
documentation for more
details concerning feature parity.
Installation
The recommended way to install this package is via PyPI:
pip install mujoco-mjx
Usage
Once installed, the package can be imported via from mujoco import mjx. Please
consult our documentation
for further detail on the package's API.
We recommend going through the tutorial notebook which introduces the MJX API
and trains a reinforcement learning policy in a few minutes:
Versioning
The major.minor.micro portion of the version number matches the version of
MuJoCo that this library provides. Optionally, if we release updates to MJX that
target the same version of MuJoCo, a .postN suffix is added, for example
3.0.1.post2 represents the second update to MJX for MuJoCo 3.0.1.
License and Disclaimer
Copyright 2023 DeepMind Technologies Limited
MuJoCo and its libraries are licensed under the Apache License, Version 2.0. You may obtain a copy of the License at https://www.apache.org/licenses/LICENSE-2.0.
This is not an officially supported Google product.