终于!这个框架重构了 OpenPI-PyTorch,与官方 JAX 版本性能对齐

具身智能之心 2026-08-04 11:00

点击下方卡片,关注“具身智能之心”公众号


你是否也遇到过这样的问题:

OpenPI 的 PyTorch 版本训练效果明显弱于 JAX 版本?

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

然而,在实际训练过程中,一个长期困扰社区的问题一直存在:

OpenPI 官方 PyTorch 实现无法稳定复现 JAX 版本性能。

这一问题尤其影响强化学习(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)。下面是其中一组实验的结果

终于!这个框架重构了 OpenPI-PyTorch,与官方 JAX 版本性能对齐图1

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

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


01.

开源代码指路



  • 代码仓库: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




02.

为什么需要

重新实现 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/dQ、dL/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`。于是问题链条就来了:

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

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

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

4. 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(...))` 对每张图独立抽样。

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

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


03.

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


面对这么多差异,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 覆盖文件的包袱。

04.

权重流通:训练到部署


模型和数据都对齐了,还剩最后一步工程收尾:权重得能在 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 高效训练 OpenPI 模型;

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


06.

RLinf介绍


RLinf是由清华大学、无问芯穹、中关村学院联合打造的首个面向具身智能大规模强化学习基础设施。

自开源以来,以开放共享为核心理念,迅速形成全球影响力。截至目前在 GitHub 已获得 4400+ Stars 、 620+ Forks、110+ Contributors,成为具身智能与大模型强化学习领域最受关注的基础设施项目之一。

RLinf 已被 Isaac Lab 官方收录为其首个面向具身大模型的训练引擎,体现了其在具身智能基础设施方向的国际技术引领能力。

此外,项目荣获 EAI-100年度十大突破奖,入选Pytorch Ecosystem、蚂蚁开源榜,代表其在具身智能基础设施领域的技术领先性与行业认可度。


END

 推荐阅读 :

终于!这个框架重构了 OpenPI-PyTorch,与官方 JAX 版本性能对齐图2


关于科技区角:国内科技展会垂直内容策划服务商,提供从论坛内容全案策划、会展市场化IP打造到精准专业观众一站式邀约服务,以产业内容吸引高质量B端人群,打通展会从议题设计、演讲嘉宾邀约、宣传预热、精准邀观到供需对接全链路。
声明:内容取材于网络,仅代表作者观点,如有内容违规问题,请联系处理。
more
首届2026中国(西部)特种电子展暨首届无人机侦测与反制对接活动成功召开!
无人机空投物资,水上航母单次渡500人,中国科技硬核救灾,这次老外真破防了……
协氢新能源:将携「L2工业级小型旋翼无人机」等,参加2026国际低空经济博览会
嘉兴政航低空运营有限公司:招聘无人机驾驶员2名,国企,base嘉兴丨低空招聘
打破严禁载人原则!近百架大疆无人机驰援广西灾区,全网看哭
比赛通知丨2026国际低空经济博览会无人机模拟飞行操控技能大赛将于7月25日召开
DJI大疆首款垂起运载无人机EV50,全球线下首次亮相!就在2026国际低空经济博览会!
获FAA第135部认证,DoorDash自研无人机配送业务“Air”正式启航
闻“汛”而动,星夜驰援!氢航氢电无人机助力广西防城港防汛抢险
UC Berkeley最新推出 | 强化学习视觉引导无人机室外飞行,零样本迁移到未知障碍环境!
Copyright © 2025-成都区角科技有限公司
蜀ICP备2025143415号-1
  
川公网安备51015602001305号