一句话定位

Google Research 用 JAX 写的开源刚体物理引擎:靠 auto-vectorization + device-parallelism + JIT + auto-diff,把成千上万个独立环境全部塞进单块加速器(GPU/TPU),让物理仿真和 RL 优化器跑在同一颗芯片上,单卡就能达到百万级 sim steps/s,几秒到几分钟训完一个 MuJoCo 式 locomotion/manipulation 策略,把 RL 训练的速度/成本改善 100–1000×;同时引擎可微,为解析策略梯度打开门。NeurIPS 2021 Datasets & Benchmarks Track。

背景与定位

论文提出的范式是 accelerator-native、on-device、可微的向量化物理仿真——把环境仿真从 CPU 搬到与学习算法同一块 GPU/TPU 上,消除跨进程/跨机的数据搬运延迟。

作者指出当时 RL 落地慢且贵的根因有三,正好对应三个设计目标:

  • 数据搬运延迟:主流引擎(mujoco、pybullet、PhysX)跑在 CPU,而 RL 算法跑在 GPU/TPU 的另一进程/另一台机器,data marshalling 与网络流量成为实验墙钟时间的主导因素。RL 样本复杂度极高——几百维状态空间的环境在探索阶段要跑百万到十亿级仿真步;业界只能靠 IMPALA/SEED-RL/OpenAI Rapid 这类大规模分布式系统硬堆,硬件与电力成本让多数研究者用不起。
  • 不可微:多数引擎是黑盒,不给环境状态的梯度,只能配 model-free RL,逼研究者用更慢、更低效的优化方法。
  • 不可内省:闭源或与 RL 技术栈完全异构,妨碍快速迭代与调试。

Brax 一次性回答这三点:物理引擎 + RL optimizer 同芯片、可微、开源且打包进 Colab(免费即可做 RL 研究)。它与同期的 NVIDIA Isaac Gym(GPU 端并行环境)属于同一波”GPU/TPU 批量化 RL”浪潮,但 Brax 走 JAX/XLA 全程 JIT + 可微 + 跨 TPU 拓扑扩展的路线。技术上继承并对话了一批可微仿真前作:Tiny Differentiable Simulator(速度级碰撞 + Baumgarte 稳定化的灵感来源)、DiffTaichi(time-of-impact 碰撞检测的取舍参照)、以及 de Avila Belbute-Peres 等的端到端可微物理。后续 Brax v2 引入 MJX(MuJoCo 的 JAX 重写)后,逐步与 mujoco 生态合流。

模型架构

这是物理引擎而非神经网络模型,“架构”指引擎的数据/计算结构。

核心状态原语 QP(maximal coordinates)

  • Brax 在 maximal coordinates(极大坐标 / 笛卡尔坐标) 下仿真:场景里每个可自由运动的独立实体单独跟踪其 position、rotational orientation、velocity、angular velocity——这就是仿真过程中唯一动态变化的数据。
  • 该数据编码为 QP(一个 flax dataclass,名字取自正则坐标 q 与 p 的戏称)。QP 带前置 batch 维:[并行场景数, 每场景刚体数, 3],例如 4 个并行场景 × 每场景 10 个刚体时 QP.pos 形状为 [4, 10, 3]——向量化因此天然。
  • joints、actuators、colliders、integration 全部实现为对这份 QP state 的 transformation(apply 函数)。如 joint.revolute 打包一个 1-DOF 约束的全部元数据,其 apply(qp) 从完整 QP 里 gather 出被约束的两个实体,返回对整份 QP 的向量化微分更新。masses、inertias、尺寸等额外数据绑定在与 QP 关联的 bodies/joints/actuators/colliders 抽象里。

物理步(Algorithm 1)

  • 一个 pseudo_physics_step(qp, action, dt):先做 kinematic 积分;再并行地对每个 joint / actuator / collider 收集 impulsive 更新(dpj / dpa / dpc);然后用 二阶 symplectic Euler(辛欧拉)积分把 joint+actuator 更新施加到 QP,最后 collision integrator 施加碰撞更新。全程尽量并行——跨 actuator、joint、collider,乃至跨整个仿真场景。
  • 顶层 system 类做协调与簿记,暴露 qp_{t+δt} = system.step(qp_t, actions),actions 即各 actuator 需要的 torque 或 target angle。扩展只需实现一个新的 Brax_transformation 并插入 step 函数。

关节与碰撞(初版实现选择)

  • 关节用 spring constraints(弹簧约束) 而非 Featherstone 式方法——大幅简化引擎核心原语,代价是需要仔细调 damping/mass/inertia scale/积分步长以保稳定,且仿真轨迹比 Featherstone 更”抖”。
  • 碰撞用 velocity-level 更新 + Baumgarte 稳定化(灵感来自 Tiny Differentiable Simulator);试过全弹性冲量碰撞但运动质量/稳定性变差。碰撞检测目前是朴素的 O(n²)(未做剪枝/宽相位),因为典型 RL 任务碰撞体不多、且可在加速器显存里并行掉。

系统规范(ProtoBuf)

  • 用 ProtoBuf 文本规范定义场景:bodies、joints、actuators、colliders(成对)。给定 body-joint 树,system.default_qp 自动求解把每个 body 放到合法关节配置的位姿。既可文本定义(类似 MuJoCo 的 XML),也可编程定义(config_pb2.Config())。默认 dt: .01

环境层

  • env 类封装 init/reset/observe/act/reward 的簿记,外加 OpenAI Gym 式 wrapper。初版内置 5 个环境(obs/act 维度):Halfcheetah 25/7、Ant 87/8、Humanoid 299/17、Grasp 139/19(四指爪 pick-and-place)、Fetch 101/10(目标导向 locomotion,玩具四足狗形态,50M frames 内可训成多种形态 locomotion),全部连续动作空间。

后续演进(GitHub README,v2):四条可互换的物理管线共享同一 API——MJX(MuJoCo 的 JAX 重写)、Generalized(广义坐标,算法接近 MuJoCo/TDS)、Positional(Position Based Dynamics,快而稳)、Spring(初版的冲量式,最快最粗)。0.13.0 起仅 brax/training 主动维护,物理仿真官方建议改用 MJX / MuJoCo Warp,环境改用 MuJoCo Playground。

数据

物理引擎无训练语料;此处的”数据规模”指仿真吞吐与内置任务集。

  • 仿真吞吐(=可生产的经验数据速率):单块现代加速器上 Ant 环境达百万级 sim steps/s;在 4×2 TPU v3 上整套环境跑到数亿 steps/s;Ant 在 8×8 TPUv3 上约”数亿 steps/s”。对照:实践者在单线程机器上跑 OpenAI Gym MuJoCo-Ant 约数千 steps/s
  • 内置任务集(初版 5 个,见上表维度):3 个 MuJoCo-like(Ant / Humanoid / Halfcheetah 的忠实但非逐位相同的重建)+ Grasp(灵巧操作 proof-of-concept)+ Fetch(目标导向 locomotion)。
  • 训练消耗的”数据量”级别示例:braxppo 扫参用 10M / 500M steps 两档;brax-sac 用 5M steps;Fetch 约 50M frames 训成 locomotion。
  • 与 MuJoCo 的偏差(Appendix E):Halfcheetah 用积分更新掩码实现平面运动而非世界关节;Ant 忽略 contact cost;Humanoid 关节 torque 正则从 MuJoCo 的 .1 降到 .01、torso done 阈值从 [1.0,2.0] 改为 [.6,2.1],3-DOF actuator 实现略异。作者明确不主张”更高 reward”,只论证 reward 随学习步的增长曲线与 MuJoCo 定性相似。

训练方法

Brax 随库提供 4 个为 JAX 并行 + JIT 定制的 RL/优化算法,全部与环境编译进同一个不中断的 jitted 函数,rollout 数据永不离开加速器:

  • PPO(on-policy):batch 均分到每个加速器核收集 rollout → 基于本 batch 算 normalization 统计并跨核同步 → 各核切 minibatch 算梯度、跨核同步后同步施加。最佳超参下瓶颈落在环境本身(Ant 上 75% 时间在跑 env)。
  • SAC(off-policy):replay buffer 完全驻留加速器,整套训练编译成单个 jitted 函数。SAC 样本效率远高于 PPO,瓶颈转到 SGD——耗时占比 12% env / 10% replay buffer / 78% SGD;因 SGD 多核扩展差,最省成本是单核。
  • ES(进化策略):lead 核生成参数扰动 → 均分到各核评估 → lead 按分数算梯度更新,>99% 时间在评估环境步。
  • APG(Analytic Policy Gradient):利用引擎可微,编译一个”对短轨迹 loss 求梯度”的函数做梯度下降,proof-of-concept;较不成熟,当前跑不出 locomotive gait、易陷局部极小(长轨迹微分是公认难题)。
  • 论文实验主打 PPO/SAC,ES/APG 留待后续。
  • 代表性超参:braxppo(Fig.3)total_env_steps 1e7、num_envs 2048、unroll_length 5、batch_size 1024、num_minibatches 16、num_update_epochs 4、lr 3e-4、discounting 0.95、entropy_cost 1e-3、reward_scaling 10、episode_length 1000。SAC(Fig.4,humanoid/ant)lr 3e-4、reward_scale 0.1、min_replay_size 1e4、num_steps 5e6、grad_updates_per_batch 64。

Infra(训练 / 推理工程)

  • 加速器与并行:全程依赖 JAX 的 vectorization(vmap)+ device parallelism(pmap)+ XLA JIT。环境计算在单卡内及跨卡分布,可扩展到”数百块互联加速器上成千上万个独立环境”。
  • 实测硬件拓扑:benchmark 用 TPUv2 / TPUv3;扩展曲线取 4×2 TPU v3,Ant 大规模跑到 8×8 TPUv3(数亿 steps/s)。Colab 免费档为 2×2 TPUv2,扫参图在 1×1 TPUv2 上生成。Fig.3 里 MuJoCo 对照跑在 128 核 Intel Xeon @2.2GHz、32×-CPU 曲线用 32 核 Xeon @2.0GHz
  • 精度:momentum/energy 守恒诊断在**单精度浮点(single precision)**下测,128 随机种子平均。
  • 速度对比:Brax 的编译优化版 PPO 在 Ant 上约 10 秒跑出可用 locomotion,标准(未编译/未并行)PPO 实现约需 半小时——两者均评估 10M 环境步。总体宣称 RL 训练速度/成本改善 100–1000×
  • 推理/控制频率:论文未给具体机器人控制 Hz/延迟/边缘硬件——Brax 是研究用仿真训练引擎,非部署栈(sim/推理均在加速器上,无 real robot HIL)。默认物理 dt = .01(100 Hz 物理步),未披露独立的策略推理时延。
  • 工程摩擦(作者自陈):JIT 编译时间对复杂环境可达数分钟、有时接近甚至超过训练时间;团队与 JAX/XLA 团队直接合作调 TPU 编译启发式来缓解,但编译时间仍是小瓶颈(对利用可微性的算法尤甚)。

评测 benchmark

一手结果均来自论文本身(多为定性曲线,作者刻意不主张”更高 reward”):

  • 加速器扩展(Fig.2):4×2 TPU v3 上五个环境的有效 env steps/s 扩展曲线;Ant 在多种加速器/TPU 拓扑下的扩展曲线(误差棒在该尺度不可见)。单卡百万级、大集群数亿级 steps/s。
  • 训练速度(Fig.3):braxppo(编译优化)vs 标准 PPO(ACME 实现),Ant,x 轴为对数墙钟秒。Brax 约 10 秒达可用 locomotion,标准实现约半小时;均 10M 步、5 seed。
  • 与 MuJoCo 的 reward 曲线对齐(Fig.4):用同一份标准 SAC(ACME,env 跑 CPU、学习在 2×2 TPUv2,非 Brax 加速版)对比 MuJoCo-{Humanoid,Ant,HalfCheetah}-v2 与 brax 对应环境,定性上相近步数达到相近 reward(HalfCheetah 有已知差距,见 Appendix E,高分策略往往是”physics-breaking”的,比较高分策略本身不严谨)。
  • 仿真质量——“astronaut”守恒诊断(Fig.5,源自 Erez et al. 2015):测线动量/角动量/能量的非守恒随仿真保真度的 scaling。Brax 凭 maximal 笛卡尔坐标 + 辛积分,线动量守恒有竞争力能量守恒与 Havok、MuJoCo 的 euler 积分器相当;角动量守恒表现尤其突出。(非 Brax 数据经原作者许可转绘;实验按 Erez 设定:动量场景关掉 damping/碰撞/重力,各 actuator 每步约 0.5 N·m 随机激励 1 秒;能量场景再关 actuator、给每个 body 1 m/s 随机初速,测 1 秒后能量漂移。)
  • 命名基线:MuJoCo(Todorov 2012)、标准 PPO/SAC 的 ACME 实现、Havok/Bullet/ODE/PhysX(守恒诊断对照)。

创新点与影响

贡献:把”物理仿真 + 学习算法”整体编译上单块加速器、全程 JIT、跨 TPU 拓扑无缝扩展,并让引擎可微——首个把这套 JAX-native 向量化可微仿真做成开源、可在免费 Colab 跑通的引擎。它证明了单卡百万级 sim steps/s、几秒到几分钟训完 locomotion/manipulation 的可行性,把原本需要大规模分布式集群的 RL 实验拉到普通研究者可及的成本区间(宣称 100–1000× 速度/成本改善)。

它改变了什么:与 Isaac Gym 一起把”GPU/TPU 上批量化数千并行环境”变成 RL 研究的默认范式,直接催生并支撑了后续 JAX RL 生态(gymnax、PureJaxRL 等)与 Brax v2 的多管线(含 MJX),最终与 MuJoCo/mujoco-playground 生态合流。QP + maximal-coordinate + transformation-as-apply 的可组合设计成为可微仿真的一种范式样板。

作者自陈的局限

  • Spring 关节脆弱、需调 damping,且质量尺度差异大时不稳;轨迹比 Featherstone 更抖;初版为速度牺牲了保真——承诺后续研究 Featherstone 方法(v2 的 Generalized/MJX 管线正是回应)。
  • 碰撞沿用 velocity-level + Baumgarte 的已知调参负担与内在非物理性;碰撞检测仍是朴素 O(n²),未做 LCP 求解器与宽相位剪枝。
  • JIT/XLA 编译时间可达分钟级,偶尔逼近训练时间。
  • APG/ES 未充分测试,长轨迹微分难优化,可微算法潜力未完全释放。
  • 社会影响:更快的控制求解引擎是双刃剑;且更快引擎可能反而诱发更多 RL 计算(如”新建高速路反增交通”),作者称实验在承诺 2030 年前全绿电的数据中心完成。

原始链接

一手源存档(sources/)

  • brax—github-readme — GitHub README 快照(v0.14.x,四管线/维护状态说明):sources/embodied/2021/brax--github-readme.md
  • arXiv 原文 PDF(2106.13281v1,全文已通读)不入 git,引用上方 arXiv URL。