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_droid、pi05_droid、pi0_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_pytorchConverting 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_droid、pi05_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),仅供参考