news 2026/9/12 2:18:09

3条命令把openpi的JAX模型搬进PyTorch

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
3条命令把openpi的JAX模型搬进PyTorch

3条命令把openpi的JAX模型搬进PyTorch

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

还在为JAX检查点进不了PyTorch生态发愁?openpi 是 Physical Intelligence 开源的机器人 VLA(视觉-语言-动作)模型仓库,它把 π₀ / π₀.₅ 的模型转换工具直接写进了 examples/convert_jax_model_to_pytorch.py:一条命令即可把 JAX 检查点导出成 PyTorch 权重,之后推理、微调、起服务全部无缝切换。读完本文你将拿到:

  • 从环境准备到权重导出的 4 步可复制命令
  • 2 处最关键的参数重排逻辑,以及为什么会这么写
  • 3 个高频坑的最短修复命令,加一段最小校验代码

整条链路很短:脚本先用 orbax(JAX 的 checkpoint 读写库)把检查点恢复成纯参数字典,再分别重排视觉塔和专家层的权重,最后实例化 PyTorch 版模型并落盘。

🚀 第一步:装好环境并确认 transformers 版本

先克隆仓库(记得带子模块),再用 uv 建环境。PyTorch 实现依赖打过补丁的 transformers,版本必须是 4.53.2。

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 . uv pip show transformers

确认版本后,把补丁文件覆盖进虚拟环境里的 transformers:

cp -r ./src/openpi/models_pytorch/transformers_replace/* .venv/lib/python3.11/site-packages/transformers/
Name: transformers Version: 4.53.2

注意:uv 默认用硬链接模式,这个覆盖会污染 uv 缓存。想彻底撤销时执行uv cache clean transformers

🔍 第二步:用 inspect_only 预览 JAX 参数结构

正式转换前,先看看检查点里有哪些参数、各是什么形状。--inspect_only只读不写,输出格式为参数名: (shape)@dtype,来自 src/openpi/training/utils.py 的array_tree_to_info

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

输出节选:

img/embedding/kernel: (14, 14, 3, 1152)@float32 llm/embedder/input_embedding: (257152, 2048)@float32 ...

官方检查点默认缓存在~/.cache/openpi下,用环境变量OPENPI_DATA_HOME可以改位置。

⚡ 第三步:3个参数一键导出 PyTorch 权重

核心就 2 个必填参数加 1 个输出目录。--config_name取自 src/openpi/training/config.py 里注册的模型名,比如pi0_droidpi05_droidpi0_aloha_sim

uv run examples/convert_jax_model_to_pytorch.py \ --checkpoint_dir ~/.cache/openpi/openpi-assets/checkpoints/pi0_droid \ --config_name pi0_droid \ --output_path ~/.cache/openpi/openpi-assets/checkpoints/pi0_droid_pytorch
Converting PI0 checkpoint from .../pi0_droid to .../pi0_droid_pytorch Model config: Pi0Config(...) Model conversion completed successfully! Model saved to .../pi0_droid_pytorch

输出目录里会得到 3 样东西:model.safetensors(全部权重)、config.json(记录 action_dim、action_horizon 等 5 个字段)、assets/(归一化统计,从检查点的上一级目录自动复制)。精度默认 bfloat16,需要全精度时加--precision float32

✅ 第四步:指向转换后目录直接起策略服务

PyTorch 权重的推理入口和 JAX 版完全一样,你只改检查点路径。起服务时把--policy.dir指向转换输出即可:

uv run scripts/serve_policy.py policy:checkpoint \ --policy.config=pi0_droid \ --policy.dir=/path/to/pi0_droid_pytorch

服务启动后监听 8000 端口等待观测数据,机器人端如何接入可看 docs/remote_inference.md。后面要用 scripts/train_pytorch.py 做 PyTorch 微调时,这个转换好的目录就是基础模型权重来源。

🧩 细节解析:两处最容易被忽略的参数处理

为什么先过一遍 JAX 加载器?检查点存储时的 dtype 和模型恢复时的 dtype 可能不一致。slice_initial_orbax_checkpoint直接调 src/openpi/models/model.py 的恢复逻辑,让权重先走一遍 JAX 模型加载,dtype 转换和真实训练时完全一致:

# examples/convert_jax_model_to_pytorch.py -> slice_initial_orbax_checkpoint params = openpi.models.model.restore_params( f"{checkpoint_dir}/params/", restore_type=np.ndarray, dtype=restore_precision )

pi05 的归一化为什么单独分支?pi0 的层归一化是只有 scale 的 RMSNorm,而 pi05 换成了带 kernel + bias 的自适应 Dense 层。slice_gemma_state_dict用检查点路径里是否含pi05字符串来区分两套参数名:

# examples/convert_jax_model_to_pytorch.py -> slice_gemma_state_dict if "pi05" in checkpoint_dir: llm_input_layernorm_bias = state_dict.pop(f"llm/layers/pre_attention_norm_{num_expert}/Dense_0/bias{suffix}") else: llm_input_layernorm = state_dict.pop(f"llm/layers/pre_attention_norm_{num_expert}/scale{suffix}")

视觉塔这边,patch 嵌入的卷积核按transpose(3, 2, 0, 1)从 JAX 的 (H, W, C_in, C_out) 转成 PyTorch 的 (C_out, C_in, H, W),这也是两个框架最常见的维度差异来源。

转换前必看的3个坑

坑 1:config_name 传了 fast 系列

  • 现象:ValueError: Config pi0_fast_droid is not a Pi0Config
  • 原因:脚本只接受 π₀ / π₀.₅ 的 Pi0Config,PyTorch 版暂不支持 π₀-FAST。
  • 修复:换成pi0_droidpi05_droid等 pi0/pi05 配置名。

坑 2:重命名了检查点目录

  • 现象:pi05 检查点转换后,归一化层权重对不上,推理行为异常。
  • 原因:脚本靠"pi05" in checkpoint_dir判断走哪套归一化分支,目录名去掉 pi05 字样就会走错分支。
  • 修复:目录名保留 pi05,如~/.cache/openpi/openpi-assets/checkpoints/pi05_droid

坑 3:输出目录缺 assets/

  • 现象:转换本身成功,但推理加载策略时报缺归一化统计的错误。
  • 原因:脚本只从检查点上一级目录assets/,找不到就静默跳过,归一化统计没跟过来。
  • 修复:转换前把源检查点的assets/放到checkpoint_dir的父目录,或转换后手动复制进输出目录。

最小校验:确认 PyTorch 权重真的能被加载

在仓库根目录跑这段代码,src/openpi/policies/policy_config.py 的create_trained_policy会通过目录里有没有model.safetensors自动切换 PyTorch 加载路径:

from openpi.training import config as _config from openpi.policies import policy_config config = _config.get_config("pi0_droid") policy = policy_config.create_trained_policy(config, "/path/to/pi0_droid_pytorch")

怎么算通过:日志出现Loading model...且没有缺键报错,说明权重被 src/openpi/models_pytorch/pi0_pytorch.py 的PI0Pytorch完整接收;再喂一组观测跑policy.infer(example)["actions"],输出第一步动作序列长度应为 10(pi0_droid的 action_horizon,pi05_droid是 15),数值落在正常关节角度范围内即视为转换成功。

收个尾

  • 转换全程 3 个参数:--checkpoint_dir--config_name--output_path,精度默认 bfloat16
  • 目录名别乱改:pi05 的分支判断依赖路径字符串,assets 也藏在检查点父目录
  • 输出目录即插即用:推理、起服务、PyTorch 微调都指向它

下一篇我们拆train_pytorch.py的单卡与多卡微调流程。有问题可以按 CONTRIBUTING.md 的指引提交 issue 或贡献补丁。

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

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

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

Flask视频播放网站开发实战:数据库、Range流与Nginx部署

简介:这是一份基于Flask框架的在线电影视频播放网站毕业设计源码,适合计算机相关专业学生用于毕设、课程设计或项目演示。网站前端采用HTML5与Bootstrap,后端使用Python3与Flask,数据库为MySQL,涵盖视频浏览、搜索筛选…

作者头像 李华
网站建设 2026/9/12 2:17:51

Node.js + AI开发实战:前端工程师快速上手大模型应用

经常有朋友私信问我类似的问题:“我不是算法工程师,也没系统学过Python,能不能做AI应用开发?”我的回答一直是:能,而且如果你本来就会一点前端或者后端,用Node.js接入AI这条路比想象中要顺得多。…

作者头像 李华
网站建设 2026/9/12 2:16:11

Apache Fesod替代EasyExcel:复杂Excel解析性能优化实战

1. 从EasyExcel切换到Apache Fesod:不是跟风,是被真实业务压出来的选择我第一次在生产环境里把EasyExcel换成Apache Fesod,不是因为看到什么技术雷达榜单,也不是听了某场分享会就热血上头——而是凌晨两点,运维同事发来…

作者头像 李华
网站建设 2026/9/12 2:14:30

STM32H750 LTDC驱动7寸RGB屏:时序参数与SDRAM显存配置全解析

简介:面向嵌入式开发者的STM32H750 LTDC驱动工程,支持7英寸1024600 RGB LCD屏,基于HAL库实现,并附带触摸屏驱动。工程覆盖LTDC控制器初始化、GPIO/时钟/DMA配置、触摸坐标解析等关键模块,源码结构清晰,便于…

作者头像 李华
网站建设 2026/9/12 2:14:06

Java魂斗罗游戏开发:帧同步渲染与实体状态机实现

简介:这是一份面向Java初学者与编程实践者的经典游戏复刻项目,基于Java SE平台实现魂斗罗核心玩法,聚焦面向对象设计、GUI绘图、事件响应、多线程控制及基础游戏逻辑构建。资源为ZIP压缩包,大小1.71MB,包含完整可运行源…

作者头像 李华