RLinf 填平 Pi 系列模型 PyTorch 性能鸿沟,为具身智能构建 SFT + RL一体化训练底座

具身智能的开发者,是不是都遇到过这样的问题:


同样是 Pi 系列模型,但 PyTorch 版本的训练效果就是追不上 JAX 版本?


由 Physical Intelligence 开源的 Pi 系列模型,是当前机器人视觉-语言-动作(Vision-Language-Action, VLA)领域最重要的基础模型之一,也是工业界和学术界广泛采用的机器人操作基线。


然而,在实际训练过程中,OpenPI 官方 PyTorch 实现却始终无法稳定复现 JAX 版本性能,已经成为具身智能行业的顽疾。

这一问题尤其影响强化学习(Reinforcement Learning, RL)后训练流程。当前机器人策略优化、在线强化学习以及大规模训练基础设施普遍依赖 PyTorch 生态,而简单地将 JAX 模型转换到 PyTorch 不仅流程复杂,还会引入性能损失,直接限制了 VLA 模型进一步通过 RL 实现持续进化的能力。


针对这一痛点,RLinf 团队对 OpenPI-PyTorch 进行了系统性重构,实现了:


PyTorch 实现与官方 JAX 版本训练性能完全对齐。


多个数据集的实验结果显示,重构后的 PyTorch 版(后文用RLinf(Pytorch)表示) 在训练 loss 曲线上与 JAX 版本(后文用OpenPI(JAX)表示)高度一致,端到端训练性能也达到同等水平,训练速度接近一致(OpenPI(JAX) 28h vs RLinf(Pytorch) 30h)。


下面是其中一组实验的结果:


在 BEHAVIOR 基准的 turn on radio 任务中,我们分别基于 OpenPI(JAX)与 RLinf(Pytorch) 进行了 SFT 训练,并保持两者使用完全相同的训练超参数。随后,我们在仿真器中进行了在线交互评估


在 128 条重复实验轨迹中,RLinf(Pytorch) 取得了 59.39% 的成功率(76/128),OpenPI(JAX) 取得了 57.03% 的成功率(73/128)。上图对比展示了两种实现训练过程中的 log losslog grad norm 指标。



关于 RLinf


RLinf 是无问芯穹联合清华大学、北京中关村学院打造的全球首个面向具身智能的大规模强化学习训练框架,始终为具身智能提供开源开放的基础设施。


自开源以来,RLinf 已在 GitHub 获得4400+Stars、630+Forks、110+ Contributors,成为具身智能与大模型强化学习领域最受关注的基础设施项目之一。


目前 RLinf 已被加州大学伯克利分校、佐治亚理工大学、香港大学、清华大学、北京大学、上海交通大学、复旦大学、中科院自动化研究所等学术机构,以及智元、自变量、原力灵机、地瓜等多家具身企业采用,并被英伟达 IsaacLab 收录为首个面向具身大模型的训练引擎,在全球学术界与产业界形成广泛的技术影响力。


与此同时,项目斩获 EAI-100年度十大突破奖,入选 Pytorch Ecosystem、蚂蚁开源榜,并亮相 2026 AI for Good 全球峰会,在具身智能基础设施领域收获广泛行业认可。


代码仓库:https://github.com/RLinf/RLinf

英文文档入口:https://rlinf.readthedocs.io/en/latest/rst_source/examples/embodied/sft_openpi_pytorch.html

中文文档入口:https://rlinf.readthedocs.io/zh-cn/latest/rst_source/examples/embodied/sft_openpi_pytorch.html



01 为什么需要重新实现 OpenPI-PyTorch?


OpenPI 的 JAX 版实现长期以来被认为是模型正确性的参考标准,而官方 PyTorch 版本为了快速适配 PyTorch 生态,主要基于 HuggingFace Transformers 组件实现,包括:

  • PaliGemma Vision-Language Backbone;

  • Gemma Action Expert;

  • SigLIP Vision Encoder。


这种方式虽然降低了开发成本,但也引入了大量隐藏差异:

  • 模型算子实现不同;

  • 数值精度路径不同;

  • Attention 计算逻辑存在差异;

  • 初始化方式不一致;

  • 数据并行方式与 JAX 不匹配。


这些差异单独看似乎只是微小数值误差,但在 VLA 大模型训练过程中会经过多层 Transformer、视觉编码器以及 flow matching head 持续放大,最终导致训练轨迹偏离 JAX 版本。


此外,官方实现使用文件覆盖的方式 monkey-patch 了 transformers 内部代码,把实现死死绑在某个 transformers 版本上,大大增加了使用包袱,使用非常别扭。


RLinf 团队系统分析了 OpenPI-PyTorch 与 JAX 实现之间的差异,共发现 22 项不一致问题,其中 11 项会直接影响训练效果。


它们大致落在三个"对不齐"的层面——模型基座侧的数值精度、模型适配侧的架构语义与初始化、数据处理侧的分片与增强。


省流版:


(一)模型基座侧的数值精度

RLinf 摒弃基于 HuggingFace 的快速拼接方案,重新实现核心模型组件:

  • 自包含实现 Gemma Transformer;

  • 自包含实现 SigLIP Vision Encoder;

  • 对齐 JAX 版本 attention 计算;

  • 修正 GELU、RoPE、RMSNorm 等关键算子数值路径;

  • 保证 float32/bfloat16 混合精度行为一致。

通过逐算子对齐,实现 PyTorch 与 JAX 前向和反向传播行为一致。


(二)模型适配侧的架构语义与初始化

OpenPI 的核心结构并不是简单的语言模型 + 动作预测头,而是:

  • PaliGemma 负责视觉和语言理解;

  • Action Expert 负责动作生成;

  • 两者通过统一 attention 机制进行交互。


RLinf 修复了原 PyTorch 实现中的 attention 交互方式,使其严格符合 JAX 版本中的:

vision-language tokens 提供条件,action tokens 基于 prefix 生成动作。

同时重新校准新增 action projection、flow matching head 等模块初始化方式,使训练过程与官方实现保持一致。


(三)数据处理侧的分片与增强

除了模型本身,RLinf 还修复了 PyTorch 分布式训练中的关键问题:

  • 修复 streaming dataset 在多 GPU 下的数据重复问题;

  • 实现正确的数据分片机制;

  • 恢复 per-image 图像增强策略;

  • 保证多机多卡训练的数据一致性。


这使 OpenPI-PyTorch 真正具备支撑大规模 SFT 与 RL 后训练的能力。



下面是详细展开版,从最隐蔽、也最容易被放过的模型基座精度问题说起。


(一)模型基座侧:继承自 HuggingFace 的精度差异


用 HF 搭建模型的好处是快,坏处是:你继承了 HF 的所有实现选择,而这些选择未必和 JAX 一致。这些选择单独看每一个都像"数值噪声",但它们同号、系统性,且在深层网络里层层累积。


先看 LLM 里最核心的三条,再补两处更隐蔽的精度差,最后是把这些毛病放大得最狠的 SigLIP。


1. GELU:tanh 近似 vs 精确


HF 的 Gemma MLP 配置为 hidden_activation="gelu_pytorch_tanh",用的是 tanh 近似:


0.5 * x * (1.0 + tanh(sqrt(2/pi) * (x + 0.044715 * x**3)))


而 JAX 的 nn.gelu 是精确的 erf-based GELU:


x * 0.5 * (1.0 + erf(x / sqrt(2)))


两者在常见区间的最大偏差不过 3e-4,看起来微不足道。但在 Gemma 的 GeGLU 结构里,GELU 的输出会直接乘到 up-projection 上(`activations = gelu(gate) * up`),误差不是加性的而是乘性的,会被 up 分支的幅度放大。更要命的是,这个激活在 18 层 Gemma FFN 和 27 层 SigLIP MLP 里各出现一次,共 45 处,误差沿残差流层层叠加。而且两个函数的导数也不同(erf 的导数是高斯,tanh 近似的导数是 sech²),所以不光前向 activation 有偏差,反向传播的梯度也在偏。训练动态从第一步起就在悄悄偏离 JAX。


2. Attention:HF eager vs JAX einsum


官方 PyTorch 走 HF 的 eager_attention_forward,内部是 scaled_dot_product_attention 那一套 reshape + scaling + mask 路径。JAX 用的是显式 einsum:


logits = jnp.einsum("BTKGH,BSKH->BKGTS", q, k, preferred_element_type=jnp.float32)



两者在数学上“应该”等价,但 QKV 的 reshape 顺序、中间累加精度、softmax 的数值实现都不同。新实现逐字翻译 JAX:logits 在 float32 下累加,softmax 之后才 cast 回 bf16。softmax 是高度非线性的,float32 与 bf16 的舍入差异会改变注意力权重的分布,而权重又直接决定梯度如何分配给 Q/K/V。所以两条路径即便前向输出接近,反传回去的 dL/dQdL/dK 也会分叉。GQA 的 BTKGH 布局尤其要小心:head 和 group 的映射顺序一旦对不上,整个 attention 就错位了。


3. RoPE:两条不同的代码路径


HF 的 GemmaRotaryEmbedding 用两步法(先预计算 cos/sin,再 apply),且在 torch.autocast(enabled=False) 上下文里强制 float32,还额外乘一个 attention_scaling 因子(Gemma 里系数是 1.0,但这条乘法路径仍然存在,会在 bf16 下引入不同的舍入)。新实现则直接计算并应用,频率与旋转全程 float32,最后一步才 .to(x.dtype) 降回 bf16,逐行对齐 JAX。RoPE 编码进的是每一层 attention 的相对位置信息,一旦有系统性偏差,token 的相对位置关系会在每一层被悄悄微调;和 attention 一样,这套舍入差异同样会进入反传,让梯度跟着偏,18 层下来整体漂移。


4. 藏在深处的两处精度差:TF32 与 AdaRMS


以上明面上的三条之外,还有两处更隐蔽:


  • TF32官方在 PI0Pytorch.init 里调了 torch.set_float32_matmul_precision("high"),等于允许所有 float32 matmul 走 TF32——尾数只有 10-bit(连符号、指数共 19-bit),而完整 float32 是 23-bit 尾数。JAX 默认不开。于是每一个 Linear、每一次 matmul 的中间精度都不一样,长训之下参数轨迹越走越偏。


  • AdaRMS:官方 monkey-patch 的 HF GemmaRMSNorm 把 non-adaptive 的 weight 显式声明成 bf16;新实现的 RMSNorm scale 默认 float32,前向内部先 x.float() 算完再 cast 回。差别在于:bf16 的 norm 权重在混合精度训练里,参数更新和梯度累积都落在低精度上,比 float32 的 scale 更容易丢掉小量更新,两者的收敛行为因此不同。


这四类差异在前向里叠加,反传时又各自变成梯度偏差——这是 LLM 路径上“看起来差不多、训练却对不上”的根因。而在视觉路径上,同样的毛病还会被放大得更狠。


5. SigLIP:同一种病,更长的放大链


SigLIP 是精度问题的重灾区——因为视觉 token 是 LLM 的输入,ViT 里的任何偏差都会被后续 18 层 Gemma 再放大一轮。上面 LLM 里的 GELU 毛病,SigLIP 又犯了一遍,还额外多出一处 stem 精度问题。


第一处是 stem 精度。官方 HF `SiglipVisionModel` 的 Conv2d patch embedding 和 position embedding 都在 bfloat16下运行,而 JAX 的 stem + pos_embed 始终 float32,算完才 cast 进 encoder。patch embedding 是整个视觉路径的第一层,bf16 的 8-bit 尾数在大量卷积乘加里累积舍入,直接改变每个 patch 的初始表示——而这正是后面 27 层的输入。新实现精确复刻 JAX,stem 和 pos_embed 都强制 float32:


x = F.conv2d(x, self.stem.weight.float(), self.stem.bias.float(), ...)
x = x + self.pos_embedding.float()
x = x.to(self.dtype_mm)          # 算完才降回 bf16 进 encoder


第二处是 FFN 的 GELU,和第 1 节同一个病:官方 SigLIP MLP 用 `gelu_fast`(tanh 近似),新实现用精确 `F.gelu`。只是这次发生在 27 层视觉 MLP 里,放大链比 LLM 更长。


把两处叠起来:初始 patch 误差 × 每层 MLP 误差 × 27 层 → 视觉特征已明显偏离 JAX → 作为 LLM 的前缀 token → 再被 18 层 Gemma 放大。视觉路径是整个模型误差累积最长的一条链。


(二)模型适配侧:结构和语义对齐


模型基座侧那些差异,本质都是"数值对不齐"。但把 PaliGemma 和 Action Expert 接进同一套 attention、再接上 flow matching head 的过程里,还有一类问题性质不同——不是精度差,而是结构或语义本身就没对上。


1. 联合 attention:统一交互 vs 手动分-合


PaliGemma 与 Action Expert 共享 attention 的正确语义是:PaliGemma 的视觉/语言 token 和 Action Expert 的动作 token 处在同一次 masked attention 里交互——动作 token attend 视觉/语言 prefix,prefix 不回看,方向完全由 mask 决定。JAX(以及新实现)的做法是把两组 token 沿序列维度 concat,统一应用 RoPE,走一次 einsum attention,再 split 回去。


官方 PyTorch 版却是手动逐层分-合:分别取 Q/K/V,拼起来做 attention,再拆回去逐 expert 走 o_proj + FFN。这条手动路径在几个边界上很脆弱——RoPE 是分段应用还是全序列统一、mask 在 concat/split 时怎么 broadcast、每段序列的 position offset 怎么维护——任何一处错位,token 交互就和 JAX 不一样了。(官方代码里甚至把 num_heads=8 硬编码进了 reshape,换个 head 数不同的模型变体就崩。)


2. 新增层的初始化


flow matching head 里的 action_in_proj / action_out_proj / 时间嵌入 MLP 是原 Gemma 没有的新层。官方用 PyTorch 默认的 Kaiming uniform 初始化,JAX 用 LeCun normal,新实现则显式 normal_(std=0.02) + zeros_(bias)。三者的方差量级并不一致:JAX 的 LeCun normal 是 std = 1/sqrt(fan_in)(随输入维度自适应),官方 Kaiming uniform 是另一套缩放,新实现固定 std=0.02。这些层只在预训练阶段需要从头训,对 SFT/ RL 的收敛没有影响。


(三)数据处理侧:两个比模型更隐蔽的坑


模型的差异再糟,也只是“慢性病”——训久了才分叉。数据处理侧有两个坑性质不同:一个让多卡训练直接失效,一个让喂进模型的数据本身就变了质。


1. MPMD 数据分片:多卡退化成单卡


第一个坑是"急性中毒":它能让多卡训练直接退化成单卡,而且不报任何错。根源在于 JAX 和 PyTorch 的分布式范式不同:


  • JAX 是 SPMD(单程序多数据):所有设备在同一个进程里,数据分片通过 `jax.sharding` 在进程内完成,每个设备天然拿到不同分片,你什么都不用做。

  • PyTorch 是 MPMD(多程序多数据,`torchrun`):每个 rank 是独立进程,各有独立的 `DataLoader`。


于是问题链条就来了:


  • 部分数据集用的是 streaming chunk 模式(例如BEHAVIOR),`getitem` 忽略传入的 `idx`,改用内部游标顺序读取。

  • 依赖 `idx` 分配数据的 `DistributedSampler` 因此完全失效

  • 当 `num_workers > 0`,DataLoader worker 用 `spawn` 启动,不继承 `torch.distributed` 状态。

  • spawn 出的子进程里 `dist.get_rank()` 永远返回 0


四步下来,所有 rank 都在复现 rank 0 的分片,每个 GPU 拿到完全相同的数据。8 卡算出 8 份相同的梯度,等价于 1 卡训练——付了 8 倍的电费,只得到 1 卡的效果。从 JAX 迁移过来的人几乎必踩这个坑,因为在 SPMD 世界里它根本不存在。


2. 图像增强:batch-level vs per-image


第二个坑是“慢性病”:数据不会崩,但分布悄悄变了质。官方 PyTorch 为了兼容 `torch.compile`,用了 tensor 级别的增强,代价是整个 batch 共享同一组随机参数——brightness、contrast、saturation 对 batch 里所有图片是同一个值。而 JAX 用 `jax.vmap(augmax.Chain(...))` 对每张图独立抽样


更糟的是官方的 contrast 归一化用的是 batch 级别的均值统计。当一个 batch 混了不同场景/任务的图像(VLA 训练里是常态),batch 均值代表不了任何一张图的真实统计,归一化后引入系统性偏移。结果是:有效增强多样性从 O(batch × steps) 掉到 O(steps),还顺带引入了错误的数据分布。新实现回到 per-image loop,逐图独立抽样、per-image 统计,对齐 JAX。


至此,模型基座侧的精度、模型适配侧的语义与初始化、数据处理侧的分片与增强,病因盘点完毕。剩下的问题是:RLinf怎么把它们一次性解决。


02 RLinf 修复原则:对齐、自包含、可验证


面对这么多差异,RLinf 的思路不是在 HF 上继续打补丁,而是自包含重写——前面每一处病因,都在这里找到对应的解法。


  • 逐行对齐 JAX。 抛弃 HF,只用 `torch` + `einops` 重新实现 Gemma(`gemma.py`)和 SigLIP(`siglip.py`),每处关键逻辑都标注对应的 JAX 源码行:精确 GELU、float32 einsum 注意力、自实现 RoPE、float32 视觉 stem、float32 RMSNorm——一一对齐。联合 attention 也回归 JAX 语义:所有 token 沿序列维度 concat,走统一 self-attention,再 split 回去。


  • 修复分布式数据。 绕开失效的 `dist.get_rank()`,改用显式的 chunk 分片:

def partition_chunk_indices(num_chunks, *, rank, world_size, worker_id, num_workers):    
global_worker_id = rank * num_workers + worker_id    
stride = world_size * num_workers    
return list(range(global_worker_id, num_chunks, stride))


rank 和 `world_size` 在主进程捕获(这时 `torch.distributed` 还有效),再通过 pickle 传给 spawn worker。`global_worker_id = rank * num_workers + worker_id` 恰好取遍 `[0, world_size × num_workers)` 里每个整数一次,而步长正好是这个总数——所以不同 `(rank, worker)` 拿到的是互斥、无重叠的等差数列,全体并集不重不漏地覆盖整个数据集。


  • per-image 增强。 回到逐图独立抽样随机参数、per-image 统计做归一化,对齐 JAX 语义。


  • 可验证。 重写之后靠四个层面兜底:单元级(固定输入,逐组件比对新实现与 JAX 的前向输出——GELU、einsum attention、`_apply_rope`、SigLIP 各层)、梯度级(backward 后比对 `dL/dQ`、`dL/dK`、FFN 权重梯度)、分布式(多 rank 各自打印样本 id,确认互斥无重叠、并集全覆盖)、checkpoint(转换后 `strict=True` 加载 + 前向比对)。


  • 自包含。不再使用 monkey-patch,去掉对 transformers 覆盖文件的包袱。


03 权重流通:训练到部署


模型和数据都对齐了,还剩最后一步工程收尾:权重得能在 JAX、旧 PyTorch、新格式、部署格式之间流转。这里的坑是:不只是改名字,还涉及真实的张量重排。SigLIP 的 Q/K/V 三个独立 Linear 要合并进 `nn.MultiheadAttention` 的 `in_proj_weight`(沿 dim=0 拼接);Gemma MLP 的 gate/up 要先 transpose 再 stack 成 `(2, in, out)`,down_proj 也要转置——忘了转置,matmul 形状要么直接崩、要么在方阵层静默算出完全错误的结果。此外官方 Action Expert 那个从没训练过的 `lm_head` 被彻底删掉了,`new2old` 反向转换时得靠 `--reference-model` 把它补回来,否则宁可报错也不产出损坏的权重。


RLinf把这些封进 5 种转换模式,共享一套核心,放在 rlinf/utils/ckpt_convertor/openpi:

  • jax2new :JAX 权重 → 新格式 

  • old2new :旧 PyTorch 权重 → 新格式

  • new2old :新格式 → 旧 PyTorch(部署兼容,需 reference model 补 action head)

  • sft2new :SFT 训练权重 → 新格式(循环剥离 FSDP/compile 等 wrapper 前缀)

  • sft2deploy :SFT 权重 → 部署格式


终点统一是 model.load_state_dict(state_dict, strict=True):任何 key 缺失或多余都立刻报错,杜绝“看起来加载成功、实际缺了几层”的静默灾难。


05 总结:为具身智能构建一体化训练基座


通过解决 OpenPI 在 PyTorch 与 JAX 实现之间长期存在的性能鸿沟,RLinf 为具身基础模型的持续训练与强化学习后训练提供统一、可靠的基础设施。此次 OpenPI-PyTorch 重构不仅是一次模型实现优化,更是 RLinf 构建 SFT + RL 一体化训练底座的重要组成部分,使研究者能够在同一框架内完成:

  • 基于 PyTorch 高效训练 Pi 系列模型;

  • 无缝衔接 SFT 与强化学习后训练流程。


站在这一成果的基础上,未来无问芯穹也将持续秉持开源开放的理念,不断打磨 RLinf 的框架易用性与工程可靠性,为物理世界的 AI 生产力打造更完善、更易用的基础设施。也欢迎各界研究者加入,共同携手推动具身智能的持续进化!



释放无穹智能,让AGI触手可及

联系我们,获取定制化 AI 基础设施解决方案

释放无穹智能,让AGI触手可及