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/scalepre_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 相关 AttributeError | transformers 没被打过补丁 | 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),仅供参考