openpi:一条命令完成 JAX 转 PyTorch,pi0 checkpoint 导出 safetensors
【免费下载链接】openpi项目地址: https://gitcode.com/GitHub_Trending/op/openpi
场景切入
openpi 的 JAX 转 PyTorch 模型转换脚本就是为这类现场准备的:仿真环境里 JAX 推理一切正常,切到产线的 PyTorch 推理服务后load_state_dict直接抛 KeyError——orbax checkpoint 的参数键名和PI0Pytorch的 state_dict 完全对不上。examples/convert_jax_model_to_pytorch.py 一次导出即可,不用逐 key 手动对维度。
一次跑通
仓库根目录跑uv sync装齐依赖:JAX、orbax、torch 2.7.1、transformers 4.53.2 都在顶层 pyproject 里锁定,PyTorch 侧无需另装。以 pi0_droid 为例,先看参数结构再执行转换:
uv sync python examples/convert_jax_model_to_pytorch.py --checkpoint_dir /home/$USER/.cache/openpi/openpi-assets/checkpoints/pi0_droid --config_name pi0_droid --inspect_only python examples/convert_jax_model_to_pytorch.py --checkpoint_dir /home/$USER/.cache/openpi/openpi-assets/checkpoints/pi0_droid --config_name pi0_droid --output_path ./pi0_droid_pytorch- 第一条只打印层级参数键树(如
llm/layers/attn/q_einsum/w),先确认 checkpoint 完整 --config_name按 checkpoint 目录名选,pi0_droid、pi05_droid、pi0_aloha_sim均可- 第二条默认以 bfloat16 写出;转换成功后输出目录有三样东西:
model.safetensors(权重)、config.json(配置)、assets/(从 checkpoint 同级目录拷入的资源)
脚本内部在做什么
错位的根源是两个框架存参数的方式不同:JAX(flax/nnx)参数是嵌套 PyTree,键为斜杠路径,Linear 的 kernel 存[in_features, out_features];PyTorch 参数是module.attr平铺 dict,nn.Linear的 weight 是[out_features, in_features]。脚本在三处做了处理:
卷积核维度转置
PyTorch 卷积类权重要求[Cout, Cin, H, W],视觉塔的 patch embedding 要做四维重排,直接照搬会 size mismatch:
# JAX 为 [H, W, Cin, Cout],PyTorch 需 [Cout, Cin, H, W] state_dict[pytorch_key] = state_dict.pop(jax_key).transpose(3, 2, 0, 1)pi05 自适应归一化层 Dense 分支
pi0 的归一化是普通 RMSNorm(一维scale参数),pi05 的动作专家改用 adaRMSNorm,归一化层带Dense线性层(kernel/bias),两代版本键名不同,需要按 checkpoint 目录名分支取参:
if "pi05" in checkpoint_dir: llm_input_layernorm_kernel = state_dict.pop(f"llm/layers/pre_attention_norm_{num_expert}/Dense_0/kernel{suffix}") else: llm_input_layernorm = state_dict.pop(f"llm/layers/pre_attention_norm_{num_expert}/scale{suffix}")MoE 多专家 state_dict 拆分与映射
PaliGemma 基座与动作专家两套权重混在同一个 dict 里,只靠键名上的_1后缀区分。脚本按expert_keys清单拆成两份映射,再分别灌入paligemma和gemma_expert两个子模块:
for key, value in state_dict.items(): if key not in expert_keys: final_state_dict[key] = torch.from_numpy(value) else: expert_dict[key] = value出错速查
🔍 三个高频报错对照:
| 报错特征 | 原因一句话 | 修复命令或参数 |
|---|---|---|
size mismatch for ... | 维度顺序错位,或 config 与 checkpoint 不对应 | 先--inspect_only核对维度,再核对--config_name |
| 推理输出明显偏离 JAX 侧 | 精度漂移 | --precision bfloat16 |
Missing key(s) in state_dict | config_name 与 checkpoint 模型版本不对应 | --config_name与 checkpoint 目录名核对,如 pi05_droid |
验证与下一步
回载校验:加载后打印动作投影层形状,与 config 的action_dim一致即权重完整——
import openpi.training.config as _config, safetensors.torch from openpi.models_pytorch.pi0_pytorch import PI0Pytorch m = PI0Pytorch(_config.get_config("pi0_droid").model) safetensors.torch.load_model(m, "pi0_droid_pytorch/model.safetensors") print(m.action_out_proj.weight.shape) # 应与 config 的 action_dim 吻合确认形状与精度后即可接入 PyTorch 推理服务,远程推理的完整部署流程见 docs/remote_inference.md,转换中遇到的其他报错可直接提交 GitHub Issue。
【免费下载链接】openpi项目地址: https://gitcode.com/GitHub_Trending/op/openpi
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考