一句话定位
MJX(MuJoCo XLA)是 DeepMind 用 JAX 重写的 mujoco 分支,随 MuJoCo **v3.0.0(2023-10-18)**首次发布,把动力学求值批量搬上 GPU/TPU(官方口径 “TPU 上可跑到百万级步/秒”),同年内两次迭代(v3.0.1 加 Newton 求解器、v3.1.0 补全 API),是 brax “generalized pipeline” 的事实继任者,也是两年后 mujoco-playground 的物理后端;2025-08(v3.3.5)又新增 NVIDIA Warp 后端(MJX-Warp)以缓解 JAX 版在接触/约束上的性能瓶颈(代价是放弃自动微分)。
背景与定位
MJX 延续的范式是 mujoco 论文奠定的”关节坐标 + 凸可逆接触模型”,只是把执行后端换成 JAX/XLA,从而获得设备端(GPU/TPU)批量并行仿真能力——这是 2021 年前后 brax、NVIDIA Isaac Gym 开启的”加速器原生 RL 仿真”浪潮里 MuJoCo 阵营的应答:把仿真器搬到与 RL 优化器同一块芯片上,消除 CPU 仿真器↔GPU 训练之间的数据搬运延迟。
MJX-JAX 官方明确点名是 brax “generalized physics pipeline” 的继任者——由同时是 MuJoCo 与 Brax 核心贡献者的团队构建;文档原话:“Brax depends on the mujoco-mjx package, and Brax’s existing generalized pipeline is no longer maintained.”(Brax 自己的 generalized 管线已停止维护,转而依赖 mujoco-mjx)。这使得 MJX 事实上统一了 MuJoCo 生态(MJCF 资产、MuJoCo Menagerie 模型库)与 Brax 的加速器批量化路线。两年后,mujoco-playground 直接建立在 MJX 之上,成为 MJX 最大规模的下游验证(覆盖四足/人形/灵巧手/机械臂并做真机 sim-to-real)。
模型架构
不是神经网络,“架构”指引擎的执行/数据结构(docs + CHANGELOG 披露):
MJX-JAX(2023-10 首发,随 v3.0.0 launch)
- 用与 C 版 MuJoCo 相同的算法重新实现前向/逆动力学,但为吃满 JAX 能力,API 故意与 MuJoCo 存在若干差异。
- 结构体:
mjx.put_model/mjx.put_data(或mjx.make_data)把mjModel/mjData拷到设备上,得到mjx.Model/mjx.Data——内含 JAX 数组,支持加 batch 维(mjx.Model加 batch 维表达域随机化,mjx.Data加 batch 维表达并行环境,天然对接jax.vmap)。部分字段是 numpy 结构字段(如jnt_limited),修改会触发 JIT 重编译;纯 JAX 数组字段(如jnt_range)运行时可改而不触发重编译。 - 函数:与 MuJoCo 同名(转 PEP8 命名),默认不 JIT,交由用户对自己的 rollout 函数整体
jax.jit。 - 求解器:CG(默认)与 Newton(v3.0.1,2023-11-15 新增,GPU 上收敛快、常仅需 1 次迭代);积分器支持 EULER / RK4 / IMPLICITFAST。
- 功能子集:仅支持 FREE/BALL/SLIDE/HINGE 关节(MuJoCo 全集的子集)、无 flex/PGS/noslip 求解器、Jacobian 仅 DENSE;可微(JAX 自动微分,“mostly supported”)。
MJX-Warp(2025-08,v3.3.5 新增,晚于本条目 2023 发布窗口,此处一并记录以说明后续演进)
- 基于 NVIDIA Warp(MuJoCo Warp / MJWarp 的封装),功能上是全集(含 flex、全部求解器、mesh collision),但不支持自动微分(官方声明”无近期支持计划”)。
- 通过
mjx.put_model(m, impl='warp')指定,需 CUDA 设备;提供graph_mode(JAX/WARP/WARP_STAGED/WARP_STAGED_EX)控制 CUDA graph 复用策略,以及硬件加速的批渲染器(create_render_context/render/get_rgb,支持pmap多 GPU)。
跨硬件支持:MJX-JAX 覆盖 XLA 支持的全部后端——Nvidia/AMD GPU、Apple Silicon、Google Cloud TPU;MJX-Warp 仅限 NVIDIA GPU(CUDA)。
数据
物理引擎,无训练语料,此栏基本不适用。文档里出现的”资产”是复用 MuJoCo 生态自身的模型:Humanoid(标准 MuJoCo 人形模型)、Barkour v0、Shadow Hand(均来自 mujoco 生态的 MuJoCo Menagerie 模型库),全部用作引擎性能基准测试的场景,而非学习数据。
训练方法
无神经网络训练;此栏解读为批量仿真与性能调优方法(docs 披露):
- 批执行:
jax.vmap对不同状态/控制并行求值,配合jax.jit编译整条 rollout;域随机化通过给mjx.Model加 batch 维实现。 - 求解器选择:Newton 在 GPU 上显著快于 CG(见下方 Infra 数值),CG 目前在 TPU 上更优。
- 性能调参项(均为官方明确建议,非默认值):调低
option/iterations/ls_iterations(RL 场景对精确解算力要求不高);显式声明contact/pair白名单减少候选接触;maxhullvert设为 ≤64 提升凸网格碰撞性能;关闭option/flag/eulerdamp;option/jacobian默认 “auto”——nv≥60或 TPU 上走 sparse,否则 dense(TPU 上 Newton+sparse 提速 2–3×,GPU 上 dense 提速 10–20%,若矩阵能放进显存);广度阶段(broadphase)用实验性max_contact_points/max_geom_pairs近似裁剪。 - GPU 环境变量:
XLA_FLAGS=--xla_gpu_triton_gemm_any=true启用 Triton GEMM emitter,在 NVIDIA GPU 上可提速约 30%。
Infra(训练 / 推理工程)
硬件:MJX-JAX 跑在 XLA 支持的全部设备(Nvidia/AMD GPU、Apple Silicon、Google Cloud TPU);MJX-Warp 仅 NVIDIA GPU(CUDA)。
吞吐(官方披露数字,均为引擎 steps/sec,非模型训练指标):
- 发布头条(v3.0.0,2023-10):“natively run MuJoCo simulations at millions of steps per second on Google TPU”(未给具体场景配置的单一数字,为定性头条声明)。
- Newton vs CG on A100(v3.0.1,2023-11-15):Humanoid 640,000→1,020,000(1.6×);Barkour v0 1,290,000→1,750,000(1.35×);Shadow Hand 215,000→270,000(1.25×)。
- CPU vs MJX-JAX GPU/TPU(“Sharp Bits”一节,1 具人形体的场景):Apple M3 Max(CPU MuJoCo,testspeed,2×numcore 线程)650K steps/s;64 核 AMD 3995WX(CPU MuJoCo,同法)1.8M;Nvidia A100 GPU(MJX-JAX,batch 8192)950K;8-chip v5 TPU(文档链接指向 Cloud TPU v5e 发布博客,MJX-JAX,batch 16384)2.7M。关键限制:随场景内人形体数(进而接触数)增多,MJX-JAX 吞吐下降比 CPU MuJoCo 更快——加速器对 broad-phase 碰撞检测里的分支代码不友好,MJX-JAX 的 broadphase 实现比 MuJoCo 简单。
- 单场景仿真:MJX-JAX 对单个
mjData实例可比(经充分优化的 CPU)MuJoCo 慢 10×;MJX-JAX 只有在批量到成千上万个并行场景时才有优势。 - MJX-Warp 吞吐(v3.3.5+,2025-08 后新增,晚于本条目发布窗口),Humanoid / Aloha Pot 场景不同 CUDA graph 模式下的 SPS:纯 Warp(无 JAX FFI)3.35M / 2.45M;JAX FFI
WARP2.96M / 2.33M;WARP_STAGED2.67M / 1.96M;WARP强制每步重新捕获图时骤降至 0.80M / 0.65M。
可微性:MJX-JAX 自动微分”mostly supported”;MJX-Warp 不支持自动微分,官方称无近期支持计划。
精度:文档未披露具体的 fp32/TF32/bf16 精度策略(仅 GPU 提到 Triton GEMM emitter 环境变量),标记未披露。
评测 benchmark
MJX 自身文档不含下游任务级 benchmark(success rate / reward 等)——它只披露引擎吞吐数字(见上方 Infra 栏的 Humanoid/Barkour v0/Shadow Hand 求解器对比、CPU-vs-加速器对比、Warp 图模式对比),且明确声明这些是仿真器性能测试而非策略性能。文档给出的唯一”效果”证据是定性的:教程 Colab 演示用 MJX + RL 在几分钟内训出人形与四足的运动策略,但未给出具体奖励曲线或成功率数字。真正的下游任务级 sim-to-real 结果(真机成功率等)出现在依赖 MJX 的后续工作(如 mujoco-playground),不在 MJX 本身的文档范围内。
创新点与影响
贡献:① 把 MuJoCo——接触丰富动力学与灵巧操作研究的事实标准引擎——原生搬到 GPU/TPU,用 JAX/XLA 而非另起炉灶,从而复用了 MuJoCo 的 MJCF 资产格式、Menagerie 模型库与既有 API 习惯;② 保留(部分)自动微分能力,同时获得设备端批量并行;③ 事实上吸收并取代了 Brax 自己的 generalized 物理管线,成为 Brax 现在依赖的后端,统一了此前 MuJoCo 与 Brax 两条并行的 “GPU 批量化 RL 仿真” 路线;④ 两年后成为 mujoco-playground 的物理底座,间接催生大批 2023–2025 年 GPU 并行 legged locomotion / manipulation sim-to-real 研究。
作者自陈的局限(docs “Sharp Bits” 一节明确列出):MJX-JAX 是 MuJoCo 功能全集的子集(关节类型受限、无 flex、无 PGS/noslip 求解器、Jacobian 仅 dense);对单场景仿真反而比 CPU MuJoCo 慢 10×,只有大批量并行才划算;大网格碰撞用分支无关的 SAT 算法,网格复杂度上升时性能与显存明显变差(建议凸分解 mesh-primitive ≤200 顶点、convex-convex <32 顶点);大规模高接触场景因加速器对分支代码不友好而吞吐下降快于 CPU 版本。这些局限严重到促使 DeepMind 在 2025-08 另起一个基于 NVIDIA Warp 的独立后端(MJX-Warp)来专门解决接触/约束瓶颈,但代价是放弃了自动微分。
原始链接
- 官方文档(MJX 页):https://mujoco.readthedocs.io/en/stable/mjx.html
- GitHub(mjx 目录,Apache 2.0):https://github.com/google-deepmind/mujoco/tree/main/mjx
- PyPI:https://pypi.org/project/mujoco-mjx/
- MuJoCo CHANGELOG(v3.0.0 首发、v3.0.1/v3.1.0/v3.1.1 迭代、v3.3.5 Warp 后端):https://github.com/google-deepmind/mujoco/blob/main/doc/changelog.rst
- 教程 Colab(人形/四足 RL 训练示例):https://colab.research.google.com/github/google-deepmind/mujoco/blob/main/mjx/tutorial.ipynb
一手源存档(sources/)
- mjx-mujoco-xla—github-readme — GitHub
mjx/README.md快照 - mjx-mujoco-xla—docs-page —
mujoco.readthedocs.io/en/stable/mjx.html官方文档快照(含 MJX-JAX/MJX-Warp 架构、feature parity、performance tuning、sharp bits 全部数字) - mjx-mujoco-xla—changelog-excerpt — MuJoCo CHANGELOG 相关版本摘录(v3.0.0/v3.0.1/v3.1.0/v3.1.1/v3.3.5)