Embodied AI Glossary中文

MuJoCo XLA

MJXCommon

A JAX implementation of MuJoCo for batched, parallel simulation on GPU/TPU that also supports differentiation.

MJX (MuJoCo XLA) is a JAX interface to MuJoCo that Google DeepMind released alongside MuJoCo 3.0 in October 2023: MuJoCo's physics are reimplemented in JAX and, once compiled through the XLA compiler, can run on NVIDIA and AMD GPUs, Apple silicon, and Google TPUs. Combined with JAX's vmap, thousands of scenes can be batched together in one computation, which suits large-scale reinforcement learning, though running a single scene alone can be roughly 10 times slower than plain MuJoCo. The pure-JAX version supports automatic differentiation, enabling differentiable simulation. MJX now has two backends: MJX-JAX, and MJX-Warp, which calls into MuJoCo Warp; the latter is faster on NVIDIA GPUs and contact-heavy scenes but cannot be differentiated. Both the Brax training library and MuJoCo Playground are built on top of it.

ExampleA typical three-step workflow: mjx.put_model loads a model onto the accelerator, mjx.make_data creates the state, and jax.vmap(mjx.step) advances thousands of robots one step at once, feeding into Brax's PPO implementation to train a walking policy.

Also called
MJX, MuJoCo MJX, mujoco-mjx, MJX-JAX
Related
MuJoCo (Multi-Joint dynamics with Contact) · JAX · MuJoCo Warp · Brax · MuJoCo Playground · Differentiable Simulation
Sources
MuJoCo XLA (MJX) documentation
mujoco-mjx (PyPI)
As of
2026-09

See it in the full glossary →