RLinf 精度对齐 + 百度百舸全链路优化:OpenPI PyTorch 训练吞吐提升 2.37 倍
作者:xxinjiang2026.08.21 15:05浏览量:65简介:RLinf 精度对齐 + 百度百舸全链路优化:OpenPI PyTorch 训练吞吐提升 2.37
经过 RLinf 的精度对齐与百度百舸的 AI Infra 全链路工程优化,OpenPI Pi0.5 在国内主流 GPU(无 NVLink、无 HPN)环境上,单机 8 卡训练吞吐从 74.20 sps 提升至 175.57 sps,达到 2.37 倍;单个 epoch 训练时间从 61.4 分钟缩短至 26.0 分钟;训练精度与基线完全一致;32 机 256 卡规模下线性扩展比保持 91% 以上,集群算力得到充分释放。
本文记录了这一 AI Infra 工程优化的全链路过程:对整条训练流水线进行性能画像,逐阶段定位瓶颈,并通过多轮迭代调优持续寻找新的最优工作点,将优化从单点改进升级为系统性工程。

具身智能正在从模型探索走向持续训练和规模化迭代。随着 VLA(Vision-Language-Action)模型不断进入 SFT、强化学习和在线交互训练阶段,训练效率正在成为影响模型研发速度的重要因素:同样一套模型和数据,训练基础设施跑得快不快,直接决定了实验迭代的周期,也影响着算力资源的实际利用率。
Physical Intelligence 开源的 Pi 系列模型,是当前 VLA 领域具有代表性的基础模型之一。OpenPI 原始实现以 JAX 为主,而当前具身智能训练与强化学习生态中,大量工程基础设施建立在 PyTorch 之上。对于希望进一步开展大规模训练和强化学习后训练的开发者而言,首先需要解决的并不是「怎么优化」,而是一个更基础的问题:如何让 Pi 系列模型进入 PyTorch 训练体系,并确保训练结果可靠。
1. 跑得对:RLinf 完成 Pi0.5 的 PyTorch 与 JAX 精度对齐
本次实践针对 OpenPI 中的 Pi0.5 模型,参数量约 3.5B,采用 MoT(Mixture-of-Transformers)架构,由三个模块组成:
三个模块串联构成完整的 VLA 前向与反向计算路径,这也是后续全链路优化中按模块拆解、分别标定的前提。
RLinf 是全球首个面向具身智能的大规模强化学习开源框架。此前,RLinf 团队已经针对 OpenPI 的 PyTorch 实现完成了与 JAX 版本的系统性对齐:共定位 22 项实现差异,其中 11 项直接影响训练效果,并从模型数值精度、架构语义、初始化以及数据处理等方面逐一修复;同时建立 forward、gradient、distributed、checkpoint 四层验证体系。
经过这轮重构,PyTorch 版本的训练 loss 曲线与 JAX 版本高度一致,端到端训练效果达到同等水平,为具身智能社区提供了一个可靠的 PyTorch 训练实现。
RLinf 官方文档地址:
https://rlinf.readthedocs.io/en/latest/rst_source/examples/embodied/sft_openpi_rlinf.html
2. 跑得快:百度百舸的全链路性能优化2. 跑得快:百度百舸的全链路性能优化
但当模型进入云上训练环境,问题又向前走了一步:跑对之后,能不能跑得足够快?
百度百舸始终立足国内现有的算力供给格局,通过系统性的 AI Infra 优化与高扩展性集群架构设计,致力于为各类具身智能场景找到高性价比的算力解决方案。在 RLinf 这一工作的基础上,我们对 Pi0.5 的整条训练链路进行了系统性的性能优化。
本次实践基于百度百舸平台 hpas.lgn7ib 实例完成。该实例搭载国内主流 GPU 型号,实例内部没有 NVLink,实例之间也未部署 HPN 网络,是国内企业和研究机构更常见的通用算力环境。最终,在核心训练配置保持一致的情况下,8 卡整机训练吞吐从 74.20 sps(samples/s) 提升至 175.57 sps,达到 2.37 倍;LIBERO 一个 epoch 的训练时间从 61.4 分钟缩短至 26.0 分钟,缩短了约 57.7%,单轮训练节省约 35.5 分钟。
但这次实践值得关注的,并不只是单机 8 卡的 2.37 倍,而是百度百舸如何从一条完整的训练流水线出发,逐阶段定位瓶颈、判断优化价值,并不断重新寻找系统的最优工作点。
2.1. 从一条训练流水线出发,而不是从一个优化点出发
一次完整的 VLA 训练,并不是一个孤立的模型计算过程,而是一条连续的数据与计算流水线:
data_fetch → H2D + 预处理 → Forward + Loss → Backward + 通信 → Optimizer
我们首先对这五个阶段进行完整 profiling,回答的不是「还能开哪些优化」,而是三个更基础的问题:这个阶段到底慢在哪里?当前瓶颈是不是主要矛盾?一个优化作用到端到端训练后,还能贡献多少收益?
基线结果很快暴露了核心矛盾:
也就是说,问题并不只是模型计算本身偏慢,更重要的是数据、计算和通信之间没有形成有效的流水。因此,这一轮优化的出发点非常明确:不是让每一个环节都平均地快一点,而是先找到真正限制端到端性能的瓶颈,再通过结构调整,让不同阶段之间形成更高效的资源协同。
需要说明的是,下文按流水线的五个阶段依次展开,这是为了叙述清晰;实际的优化实施顺序并不等于流水线顺序,而是由每一轮 profiling 的结论决定:先做收益最确定的 slot 剪枝与通信结构重构,等通信结构改变、显存大量释放之后,再回头重新标定前向阶段的重计算策略。
2.2. data_fetch:先减少真正需要处理的数据
VLA 训练天然包含大量视觉数据,因此,数据进入模型之前发生了什么,会直接决定后面需要承担多少计算。
本次 LIBERO 训练使用 LeRobot v2.0 数据集(1693 个 episodes、273465 帧、40 个任务、动作维度 7),实际数据只有两路相机;而 Pi0.5 接口预留了三个 image slots。原始流程会将第三路图像填充为全零数据。虽然这一位置最终会通过 Attention Mask 被屏蔽,但在这之前,它仍然完整经历 H2D、SigLIP 视觉编码以及后续 Attention,相当于 GPU 为一张最终不会产生有效结果的图像完成了完整计算。
我们没有选择继续优化这部分无效计算,而是把 slot 数量做成自适应的:slot 数不是从配置里读出的固定值,而是每个 batch 在 collate 之前根据数据本身「感知」出来,无效 slot 直接剪掉,只让真正有效的数据进入模型。这也意味着同一套逻辑可以直接适配不同相机路数的数据集。这样,视觉输入从 3 个 slot 减少到 2 个,prefix token 从 968 个下降到 712 个,Attention score 矩阵的元素数量约降低到原来的 54%。
最终,这一项优化单独带来了 +27.13 sps 的吞吐提升,是整个优化链中收益最大的单项之一。
2.3. H2D + 预处理:一项收益有限但零成本的优化
数据进入 GPU 后,我们继续检查 H2D 与预处理路径。
其中一个优化是将图像由 float32 形式传输改为 uint8 H2D,再在 GPU 上完成归一化。这样可以把 H2D 数据量降低到原来的四分之一,实测数据传输量从 29.0 MB 降至 9.7 MB,并通过逐位 parity 测试验证结果 bit-exact。
但有意思的是,这项优化最终只带来了 +1.32 sps 的端到端吞吐提升。
如果只看局部数据传输量,4× 的下降看起来非常可观;但放回完整训练链路之后,它并不是当前工作点的主要瓶颈,因此端到端收益非常有限。
这项优化零成本且结果 bit-exact,因此保留;工程投入则继续集中到主要瓶颈上。
2.4. Forward + Loss:既要减少执行开销,也要重新分配「计算」和「显存」
进入 Forward 阶段后,问题从数据流转向模型计算。
Pi0.5 中的 18 个 Gemma Block 在 eager 模式下会产生大量独立的 GEMM、pointwise 和 reduction kernel。对于 prefix 和 action suffix 等形状相对固定的训练路径,这类计算非常适合通过编译进行融合。
我们使用 torch.compile + Inductor 对 Gemma Block 进行编译优化,减少 kernel 发射以及碎片化计算带来的开销,最终带来 +15.88 sps 的吞吐提升。
但优化到这里,新的问题又出现了:计算和显存究竟应该如何分配?
基线对 45 个 Transformer Block 全部进行 activation checkpointing,本质上是通过「多计算一些」来换取「少占一些显存」。然而,随着后续通信结构重构(FSDP,详见「Backward + 通信」一节)释放大量显存,原来的重计算策略已经不再是最优工作点。如果继续维持原有策略,就相当于在拥有更多显存之后,仍然承担不必要的重复计算。
因此,我们重新对 activation recomputation 进行标定。第一次调整后,吞吐提升 11.15 sps;随后进一步对 Gemma 和 SigLIP 分别进行联合标定,又获得 +9.48 sps 的收益。
这里出现了一个很典型的工程判断:当 SigLIP 的重计算 stride 从 1 调整到 4 后,显存占用增加了约 7 GiB,但速度仅提升 0.09%。继续「放开」重计算,并没有带来对应的收益。原因在于 SigLIP 侧的激活张量更大,而重计算本身代价相对较低,释放重计算所换来的计算收益并不值得支付额外显存。
因此,最终采用的是不同模块分别标定的策略:Gemma 更值得用显存换速度,而 SigLIP 则保持更高的重计算比例。
2.5. Backward + 通信:从「通信存在」到「让通信被计算覆盖」
如果说前面的优化主要围绕数据和计算展开,那么进入 Backward + 通信阶段之后,我们面对的是这次实践中最核心的系统瓶颈。
基线情况下,3.5B 模型几乎退化成一个大的 FSDP 单元。FSDP(Fully Sharded Data Parallel)是分布式训练中用于切分模型参数的机制,通过将参数、梯度和优化器状态分布到各 GPU 上,以降低单卡显存压力。但在此基线中,参数在 Forward 开始时集中进行 AllGather,梯度在 Backward 阶段集中进行 ReduceScatter,通信和计算之间没有形成有效重叠,因此 NCCL 大量暴露在关键路径上。
百度百舸团队首先重新设计 FSDP wrap,将原来的大单元拆成 block 级单元,使 Gemma 与 SigLIP 具备更细粒度的参数管理和通信能力。
单看吞吐,这一步只有 +3.67 sps,并不起眼;但它同时把显存占用从 68.37 GiB 降到 32.45 GiB,直接释放 35.92 GiB 显存。更重要的是,这次结构变化让训练系统第一次拥有了足够细粒度、可以被提前调度和预取的通信对象。
因此,Per-block FSDP 不能只看自身的吞吐增益。它改变的是后续优化的空间。
在此基础上,我们进一步开启 forward prefetch 和 backward prefetch,让下一阶段的通信提前发生,并尽可能隐藏在当前计算之后。结果非常明显:
Prefetch 单项带来了 +28.25 sps 的吞吐提升,成为整条优化链中按吞吐口径收益最大的单项(按步时降幅,slot 剪枝仍为第一,两个口径不矛盾)。
这也是本次实践最典型的系统优化案例。同样一个 Prefetch,如果直接作用在原来的 root-only FSDP 上,因为缺少足够细粒度的通信对象,几乎没有收益;而当 FSDP 结构被拆细之后,它才找到了自己的「作用对象」。
因此,性能优化并不是简单的开关叠加,这里存在明显的前置关系:
FSDP 重构 → 获得细粒度通信单元 → Prefetch 有了可调度对象 → 通信进入计算间隙 → NCCL 裸露时间下降 → 整体吞吐提升。
在 FSDP 拆细和 Prefetch 生效的基础上,团队进一步对分组粒度做了调整。Gemma 层保持 1 层一组,SigLIP 层合并为 4 层一组,AllGather 次数从 92 次降至 52 次,带来 +2.77 sps 的吞吐提升,同时将步时抖动压至全链最低。
2.6. Optimizer:边际收益与为后续让路的优化
完成主要的计算与通信优化之后,我们继续检查 Optimizer 和运行时环节。
例如,fused AdamW 将原本分散的参数更新过程进一步融合,带来约 +0.76 sps 的吞吐提升;显存分配器碎片优化(expandable_segments)自身的吞吐收益只有 +0.44 sps,但可以回收约 9.78 GiB 显存。
因此,在完整训练系统里,优化项承担的角色并不完全相同:
还有一类优化,价值会随工作点变化。循环 GC 对齐在较早的 mbs16(micro batch size 16) 工作点上曾有效消除各卡随机 GC 带来的步时抖动;但在本次 mbs32 工作点上,步时抖动本已很小,收益落入噪声范围,周期性触发反而引入了新的尖峰。团队最终没有将其作为性能优化项,仅保留为可复现配置。
此外,重算随机状态、数据加载预取深度等零成本细项,端到端收益均在噪声范围内,作为可复现配置保留。
3. 真正难的不是做 13 个优化,而是不断寻找新的最优工作点
如果把这次实践简单理解成「完成了 13 项优化」,其实会低估其中的工作量。因为这些优化从来不是彼此独立的。
Per-block FSDP 自身只带来有限的吞吐提升,却释放了 35.92 GiB 显存,为后续 activation recomputation 的重新标定提供了空间;与此同时,它又改变了 FSDP 的通信结构,使 Prefetch 从「几乎没有作用」变成整条链路中收益最大的优化之一。显存分配器碎片优化也是如此:自身吞吐贡献很小,但回收出来的显存继续被投入后续优化。甚至 activation recomputation 都不能脱离 FSDP 单独判断:在不同的 FSDP 工作点下,同一个重计算 stride,其显存成本可以相差 12 倍。
这意味着,系统优化的过程更像一个不断迭代的闭环:
profiling → 找到瓶颈 → 改变结构 → 获得新的资源空间 → 重新标定 → 再次 profiling。
每完成一轮优化,系统的工作点就发生一次变化;而工作点发生变化,原来的瓶颈、优化收益和资源边界也会随之变化。
所以,最终的 2.37 倍不能通过简单相加或者相乘得到。原始实践采用的是同一条完整链路上的逐步累加,在统一的训练配置下持续验证每一项优化的真实端到端贡献。
4. 结果
经过完整优化,在 mbs32 × 8 卡 × accumulation 1 的统一配置下:
这意味着这次优化带来的并不只是某几个 kernel 更快,而是训练系统的执行方式发生了改变:原本暴露在关键路径上的通信,被逐步隐藏到计算背后;原本不必要的视觉计算被直接删掉;原本用于换取显存的重复计算,则重新根据工作点进行了分配。
5. 精度与扩展性:跑得快之外的两个前提
吞吐提升的前提是精度不变。在 Pi0.5 的 8 卡 1000 步对照训练中,优化后版本与基线的学习率逐位相同,全程平均 loss 仅相差 0.0001,训练末期每步 loss 均值偏差小于 0.001,两条训练曲线高度一致。2.37 倍的吞吐提升,没有以任何精度损失为代价。
单机性能达到上限之后,训练还需要能继续扩展。在上述优化的基础上,团队近期又完成了新一轮调优,扩展性测试即基于这一最新版本,单机基准吞吐达到 182.99 sps,略高于前文的 175.57 sps。在实例间未部署 HPN 网络的条件下,依托百度百舸的 ERI(弹性 RDMA 互联)网络能力,训练规模从 1 台实例扩展至 32 台实例(256 卡):


32 台实例的整机吞吐达到 5340 sps,线性扩展比保持在 91% 以上(32 台为 91.2%)。这意味着前文的全部单机优化在多机场景下依然成立:通信与计算的重叠结构没有因为跨实例通信的加入而失效。
6. 在百度百舸,一键使用这些优化能力
上述全部优化已沉淀为百度百舸平台的预置镜像。在百舸控制台「快速开始」中选择 RLinf v0.3 加速版镜像,即可一键拉起预装加速运行环境的开发机,内置 LIBERO 数据集与 pi05_base 基座权重,支持单机 8 卡与多机分布式训练,开箱即可复现本文的训练性能。
7. 从「跑得对」到「跑得快」,背后是 AI Infra 的全链路能力
回到最开始的问题,RLinf 与百度百舸解决的是两个不同层次的问题。
RLinf 解决的是 Pi0.5 在 PyTorch 生态中能不能跑对的问题。它完成了从 0 到 1 的实现拓荒,让 Pi 系列模型能够在 PyTorch 生态中稳定开展 SFT 与 RL 训练。
百度百舸进一步解决的是跑对之后能不能跑快的问题:从数据、计算、通信、显存到运行时,对一次真实训练任务进行系统的性能分析,并在每一轮变化后的工作点上,重新判断每项优化的价值。
这次单机 8 卡的 2.37 倍,以及 256 卡规模上 91% 以上的扩展效率,正是这套全链路方法的结果。
对于正在进入规模化训练的具身智能而言,跑得对,是进入训练的前提;跑得快,决定模型持续进化的速度。
而百度百舸希望解决的,正是这条从模型到基础设施、从一次训练到持续迭代的效率问题。
感谢 RLinf 团队扎实的工程贡献。他们完成的 PyTorch 与 JAX 精度对齐是本次优化的起点和前提,让 Pi0.5 模型能够在 PyTorch 生态中稳定运行。

登录后可评论,请前往 登录 或 注册