news 2026/9/12 8:30:58

3 步命令完成 JAX→PyTorch 模型转换:openpi 实战手册

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
3 步命令完成 JAX→PyTorch 模型转换:openpi 实战手册

3 步命令完成 JAX→PyTorch 模型转换:openpi 实战手册

【免费下载链接】openpi项目地址: https://gitcode.com/GitHub_Trending/op/openpi

你的微调流水线跑在 JAX 上,线上推理栈却全是 PyTorch,这个断点卡住过不少 VLA 开发者。本文将带你用 openpi 的官方脚本完成 JAX→PyTorch 模型转换:2 条命令转换检查点、4 项排障速查、外加一段 9 行验证代码,让你判断转换后的权重是否可用。

项目速览:openpi 定位与转换支持边界

openpi 是 Physical Intelligence 开源的机器人模型仓库,包含 π₀(flow matching VLA)、π₀-FAST(自回归 VLA)、π₀.₅(开放世界泛化增强版)三类模型及配套训练/推理工具链。它的模型转换工具覆盖 pi0 与 pi05 两个家族;π₀-FAST、混合精度训练、FSDP、LoRA 目前不在 PyTorch 版支持范围内(详见 README 的 PyTorch Support 一节),动手前先确认你的目标模型在支持清单内。PyTorch 实现已在 LIBERO 基准上验证过推理与微调,配合torch.compile后推理速度与 JAX 相当。

图中 C、D、E 节点分别对应 slice_paligemma_state_dict、slice_gemma_state_dict 与 convert_pi0_checkpoint 内部的 projection 参数处理段,B 节点入口是 restore_params。

动手实操 🔧:3 条命令完成转换

第 1 步:环境初始化并打 transformers 补丁

克隆仓库、用 uv 装依赖,再把打过补丁的 transformers 覆盖到本地环境(PyTorch 实现依赖 AdaRMS 等扩展行为)。

git clone --recurse-submodules https://gitcode.com/GitHub_Trending/op/openpi cd openpi && GIT_LFS_SKIP_SMUDGE=1 uv sync && GIT_LFS_SKIP_SMUDGE=1 uv pip install -e . cp -r src/openpi/models_pytorch/transformers_replace/* .venv/lib/python3.11/site-packages/transformers/

看什么:uv pip show transformers输出版本号为4.53.2——补丁就是针对这个版本打的,版本不符后面会出怪错。另注意 uv 默认 hardlink 模式下,这次覆盖会连带改动 uv 缓存,想彻底撤销需执行uv cache clean transformers

第 2 步:转换前查看 JAX 参数层级

先 dry-run 一遍检查点,确认路径可用、参数命名符合预期。

uv run examples/convert_jax_model_to_pytorch.py \ --checkpoint_dir ~/.cache/openpi/openpi-assets/checkpoints/pi0_droid --config_name pi0_droid --inspect_only

看什么:控制台打印分层参数键树(附带 shape/dtype 信息),形如:

img/embedding/kernel img/pos_embedding llm/layers/attn/q_einsum/w llm/layers/mlp_1/gating_einsum llm/final_norm_1/scale

pre_attention_norm_1下若挂的是Dense_0/kernel,说明这是 pi05 家族的自适应归一化参数;只有裸scale则是标准 pi0——这一眼就能看出转换脚本将要走哪个分支。另外--config_name是位置参数,只查看时也不能省。

第 3 步:执行模型转换,产物与参数一次说清

指定--config_name与输出目录,精度默认bfloat16

uv run examples/convert_jax_model_to_pytorch.py \ --checkpoint_dir ~/.cache/openpi/openpi-assets/checkpoints/pi0_droid \ --config_name pi0_droid --output_path ./pi0_droid_pytorch --precision bfloat16

看什么:控制台依次打印:

Converting PI0 checkpoint from .../pi0_droid to ./pi0_droid_pytorch Model conversion completed successfully! Model saved to ./pi0_droid_pytorch

输出目录产物清单:

  • model.safetensors:PyTorch 权重,下游load_pytorch只认这一个文件
  • config.json:记录action_dim/action_horizon/precision等字段,供参考
  • assets/:仅当原检查点同级的上级目录存在 assets 时才会被原样拷贝,属附加资源

原理拆解:权重对齐的三个关键点

关键点一:NHWC→NCHW 决定了卷积权重必须先转置

JAX 的卷积核按 [H, W, C_in, C_out] 存储(对应 NHWC 布局),PyTorch 的Conv2d权重形状则是 [C_out, C_in, H, W]。把 JAX 数组直接塞给 PyTorch 会触发 size mismatch,或更糟——维度恰好能对上但语义错位,属于静默广播类错误,patch embedding 必须先转置:

# JAX 按 NHWC 存核为 [H,W,Cin,Cout],PyTorch 要 [Cout,Cin,H,W] state_dict[pytorch_key] = state_dict.pop(jax_key).transpose(3, 2, 0, 1) # Gemma 的 einsum 注意力权重是 3D 张量 (层, 头, 维),先转维再拍平成标准 Linear q = llm_attention_q_einsum[i].transpose(0, 2, 1).reshape(num_heads * head_dim, hidden_size)

这样设计是因为 JAX/XLA 沿 NHWC 优化,cuDNN 沿 NCHW 优化,两套布局互为镜像,转置是唯一正确的对齐方式。同理,Gemma 用 einsum 把每个注意力头的权重独立存成三维张量,不拍平成二维矩阵就无法被nn.Linear接收。

关键点二:pi05 与 pi0 的 LayerNorm 结构不同,脚本靠路径自动判别

π₀.₅ 引入了自适应归一化(AdaRMS),每层 norm 是一个带 kernel/bias 的Dense_0线性层;π₀ 系是标准 RMSNorm,只有一维scale向量。转换脚本通过检查点路径决定走哪个分支:

# pi05:AdaRMS,Dense_0 带 kernel/bias if "pi05" in checkpoint_dir: kernel = state_dict.pop(f"llm/layers/pre_attention_norm_{num_expert}/Dense_0/kernel{suffix}") else: # pi0:标准 RMSNorm,只有 scale 向量 scale = state_dict.pop(f"llm/layers/pre_attention_norm_{num_expert}/scale{suffix}")

脚本作者在这里选了最轻量的判别方式:不用扫整棵参数树推断模型家族,直接看路径里有没有pi05字样。代价是检查点目录名必须有意义(命名里带 pi05 才走对分支),这也解释了为什么--config_name配错是最常见的坑。

关键点三:先按 float32 恢复,最后一步才降精度

orbax 检查点按 bfloat16 存储,但转换流水线第一步就强制按float32恢复成 numpy:

# orbax 存的是 bfloat16,先按 float32 恢复, # 转完再统一 cast,避免中间过程累积精度损失 initial_params = slice_initial_orbax_checkpoint(checkpoint_dir, restore_precision="float32")

原因是中途的 transpose/reshape 不改变数值,只有精度转换会,先高恢复、转换结束后按--precision一次 cast(见 convert_pi0_checkpoint 的 L522-L527),最终 safetensors 的精度就由这一个参数唯一决定。顺带提醒:--precision虽然声明了Literal["float32", "bfloat16", "float16"],但实现里只接受前两个,传float16会直接抛ValueError

排障速查:高频 4 类报错与修复

症状根因修复动作
Error: --output_path is required转换模式没给输出目录--output_path <dir>;只想看参数就加--inspect_only
ValueError: Config xxx is not a Pi0Config--config_name指向了非 pi0 系配置(如 pi0_fast 系)换成 pi0/pi05 家族的 config;PyTorch 版暂不支持 π₀-FAST
size mismatch/Missing key(s)--config_name与检查点版本不匹配,pi05 与 pi0 的 AdaRMS/RMSNorm 结构对不上让 config 名与检查点版本一致,pi05 检查点就用 pi05 系 config
加载时报 AdaRMS 相关 AttributeErrortransformers 没被打过补丁uv pip show transformers确认 4.53.2,再执行第 1 步的cp -r覆盖

最容易踩的是最后一行:PyTorch 实现依赖覆盖 transformers 文件带来的三处行为——支持 AdaRMS、正确控制激活精度、KV cache 可在未更新时使用(README PyTorch Support 一节有解释)。先确认uv pip show transformers输出是 4.53.2,再重跑第 1 步的cp -r命令;uv hardlink 模式下补丁会渗透进缓存,彻底回滚靠uv cache clean transformers

落地验证 ⚡:9 行代码确认权重可用

最直接的验收方式:实例化PI0Pytorch、加载转换出的 safetensors,检查精度与配置落盘是否正确。

import json, safetensors.torch from openpi.models_pytorch import pi0_pytorch from openpi.training import config as _config model_config = _config.get_config("pi0_droid").model model = pi0_pytorch.PI0Pytorch(model_config) safetensors.torch.load_model(model, "./pi0_droid_pytorch/model.safetensors") print(next(model.parameters()).dtype) print(json.load(open("./pi0_droid_pytorch/config.json")))

预期输出(具体维度值随检查点而定):

torch.bfloat16 {'action_dim': …, 'action_horizon': …, 'paligemma_variant': …, 'action_expert_variant': 'gemma_300m', 'precision': 'bfloat16'}

第一行是torch.bfloat16、第二行的precision字段与你传参一致,即转换闭环成立。往下接推理也很轻:policy_config.create_trained_policy会凭检查点目录里是否存在model.safetensors自动切到 PyTorch 路径(见 policy_config.py),policy.infer(example)的 API 与 JAX 版完全一致。

接下来可以做什么

从 orbax 检查点到可部署 safetensors,openpi 把 JAX→PyTorch 的迁移压缩成了一条命令。下一步可以直接用scripts/train_pytorch.py在转换后的权重上做继续微调(torchrun支持单节点多卡);想往边缘端走的话,量化与蒸馏是官方 PyTorch 实现尚未覆盖的空档,也是不错的贡献切入点——踩到本文没收录的坑,欢迎在仓库提 Issue 或按 CONTRIBUTING.md 补一份文档。

【免费下载链接】openpi项目地址: https://gitcode.com/GitHub_Trending/op/openpi

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/12 8:30:55

用 DolphinScheduler 把一条数仓同步链路调度起来

用 DolphinScheduler 把一条数仓同步链路调度起来 【免费下载链接】dolphinscheduler Apache DolphinScheduler is the modern data orchestration platform. Agile to create high performance workflow with low-code 项目地址: https://gitcode.com/GitHub_Trending/dol/d…

作者头像 李华
网站建设 2026/9/12 8:29:10

三相逆变器参数设计实战:从公式到硬件的工程闭环

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/12 8:25:05

SpringBoot+Netty实现物联网高并发通信方案

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/12 8:24:42

SVG动画可交付性:从AI生成到生产上线的工程实践

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/12 8:23:48

构建高效测试体系:从方法论到自动化实践

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华