一句话定位

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_modeJAX/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 v0Shadow 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/eulerdampoption/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 WARP 2.96M / 2.33M;WARP_STAGED 2.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)来专门解决接触/约束瓶颈,但代价是放弃了自动微分。

原始链接

一手源存档(sources/)