
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_SMUDGE1 uv sync GIT_LFS_SKIP_SMUDGE1 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.safetensorsPyTorch 权重下游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向量。转换脚本通过检查点路径决定走哪个分支# pi05AdaRMSDense_0 带 kernel/bias if pi05 in checkpoint_dir: kernel state_dict.pop(fllm/layers/pre_attention_norm_{num_expert}/Dense_0/kernel{suffix}) else: # pi0标准 RMSNorm只有 scale 向量 scale state_dict.pop(fllm/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_precisionfloat32)原因是中途的 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_onlyValueError: Config xxx is not a Pi0Config--config_name指向了非 pi0 系配置如 pi0_fast 系换成 pi0/pi05 家族的 configPyTorch 版暂不支持 π₀-FASTsize 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.pypolicy.infer(example)的 API 与 JAX 版完全一致。接下来可以做什么从 orbax 检查点到可部署 safetensorsopenpi 把 JAX→PyTorch 的迁移压缩成了一条命令。下一步可以直接用scripts/train_pytorch.py在转换后的权重上做继续微调torchrun支持单节点多卡想往边缘端走的话量化与蒸馏是官方 PyTorch 实现尚未覆盖的空档也是不错的贡献切入点——踩到本文没收录的坑欢迎在仓库提 Issue 或按 CONTRIBUTING.md 补一份文档。【免费下载链接】openpi项目地址: https://gitcode.com/GitHub_Trending/op/openpi创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考