2 步完成 JAX 转 PyTorch 权重转换:openpi 检查点迁移避坑全记录
2026/9/12 8:40:53 网站建设 项目流程

2 步完成 JAX 转 PyTorch 权重转换:openpi 检查点迁移避坑全记录

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

JAX 训练出来的 pi0 检查点,直接喂给 PyTorch 部署栈就是一串size mismatch。用 openpi 的 JAX 转 PyTorch 权重转换脚本,2 步、6 条命令,把 Orbax 检查点变成 PyTorch 能直接加载的 safetensors。

JAX 转 PyTorch 的第一道坎:Orbax 权重布局 PyTorch 不认

想拿 pi0 / pi05 检查点跑 PyTorch 推理或微调的部署同学,卡就卡在权重格式上:检查点是 Orbax 存的 JAX 参数,卷积核和 einsum 注意力的布局,nn.Linear根本不认。转换工具在整条链路里的位置:

下面直接进命令。

快速上手:2 步完成检查点转换

1. 装好代码和依赖

git clone --recurse-submodules https://gitcode.com/GitHub_Trending/op/openpi cd openpi GIT_LFS_SKIP_SMUDGE=1 uv sync # 拉全部依赖(LeRobot 需跳过 LFS) GIT_LFS_SKIP_SMUDGE=1 uv pip install -e .
# 预期输出(N 视环境而定,无报错即成功) Resolved N packages Installed N packages

卡住了?最高频的失败是uv sync报依赖冲突——多数是 LeRobot 子模块没拉下来。用git submodule status排查,缺了就补git submodule update --init --recursive

2. 先检查参数键,再落盘转换

# 第 1 步:只检查参数键结构,不落盘(转换前务必先跑) uv run examples/convert_jax_model_to_pytorch.py \ --checkpoint_dir ~/.cache/openpi/openpi-assets/checkpoints/pi0_droid \ --inspect_only # 第 2 步:执行转换(pi05 系列把 --config_name 换成 pi05_droid) 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 ... to ... Model config: Pi0Config(...) Model conversion completed successfully! Model saved to ~/.cache/openpi/openpi-assets/checkpoints/pi0_droid_pytorch

卡住了?最高频的是checkpoint_dir不存在。检查点首次使用会自动从gs://openpi-assets下载到~/.cache/openpi缓存(可用OPENPI_DATA_HOME改位置),先ls ~/.cache/openpi/openpi-assets/checkpoints确认目录在不在。

转换成功后的产物清单:

  • model.safetensors:转换后的 PyTorch 权重,create_trained_policy靠它自动识别走 PyTorch 分支
  • config.json:记录action_dimaction_horizonpaligemma_variant、精度等,供人工核对
  • assets/:从检查点同级目录拷贝,含推理必需的norm_stats.json

关键机制拆解:维度对齐与专家权重拆分

slice_paligemma_state_dict:视觉塔与 LLM 权重的维度对齐

它解决的具体问题是:把 JAX 的卷积核和 einsum 矩阵,翻译成 PyTorchLinear认的二维布局。

# examples/convert_jax_model_to_pytorch.py 第 57、188-194 行附近 state_dict[pytorch_key] = state_dict.pop(jax_key).transpose(3, 2, 0, 1) # 卷积核换轴 q_proj_weight_reshaped = ( llm_attention_q_einsum[i] .transpose(0, 2, 1) # einsum 轴序重排 .reshape(heads * head_dim, hidden_size) # 拼成 [out, in] ) state_dict[f"...layers.{i}.self_attn.q_proj.weight"] = q_proj_weight_reshaped

如果跳过这步直接load_state_dict,q/k/v 全部 size mismatch,模型一步都跑不起来。

slice_gemma_state_dict:pi0 与 pi05 的归一化分支

动作专家权重和主干 LLM 挤在同一个 dict 里,脚本按前缀拆层;而 pi05 的自适应归一化不再是 scale 向量,必须走不同分支。函数用if "pi05" in checkpoint_dir:(第 293 行附近)区分:pi05 把pre_attention_norm_1/Dense_0/kernel|bias装进input_layernorm.dense.weight/bias,否则把scale装进input_layernorm.weight

如果 pi0 和 pi05 检查点混着用,报Missing key(s)或 shape 对不上,八成就是这条分支走错了。

slice_initial_orbax_checkpoint:为什么要绕 JAX 加载器读权重

同一份检查点在不同训练配置下 dtype 不同,直接读 Orbax 原始文件会拿到错的精度。脚本用restore_params(f"{checkpoint_dir}/params/", restore_type=np.ndarray, dtype="float32")(第 401 行附近)走 JAX 模型的 restore 路径,让 dtype 转换和 JAX 训练时一致。绕过它直接读原始 shard,混合精度检查点转换后的推理数值会整体漂移。

踩坑对照表

⚠️ 只收录实际跑转换时的高频报错,不凑数:

现象(报错关键字)根因(一句话)修复(命令或代码,≤ 2 行)
Invalid precision: float16脚本只落地支持 float32 / bfloat16,float16 走 else 抛错--precision bfloat16(默认值,可省略)
Error: --output_path is required没传--inspect_only却忘了输出路径--output_path <dir>
Config xxx is not a Pi0Config--config_name指到了 pi0_fast 等非 flow 版本换成 pi0 / pi05 系列,如pi0_droid
Missing key(s)/size mismatchpi0 与 pi05 归一化层结构不同,检查点和 config 版本没对上--config_name与版本成对:pi0 对pi0_droid,pi05 对pi05_droid

转换前先用--inspect_only跑一遍,键结构没问题再落盘,省一次完整转换的时间。--config_name和检查点版本(pi0 还是 pi05)必须成对,这是绝大多数报错的源头。

验证 & 下一步

from openpi.training import config as _config from openpi.policies import policy_config config = _config.get_config("pi0_droid") # 与转换时的 --config_name 一致 policy = policy_config.create_trained_policy( # 目录里有 model.safetensors 自动走 PyTorch config, "~/.cache/openpi/openpi-assets/checkpoints/pi0_droid_pytorch" ) actions = policy.infer(example)["actions"] # example 为观测 dict,键见 README print(actions.shape) # shape 与 JAX 版输出一致即通过 ✅

推理服务与远端部署对接见docs/remote_inference.md,想给转换脚本支持新模型从CONTRIBUTING.md入手。下一篇看scripts/train_pytorch.py怎么在 PyTorch 下微调 pi0。

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

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

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询