
2 步完成 JAX 转 PyTorch 权重转换openpi 检查点迁移避坑全记录【免费下载链接】openpi项目地址: https://gitcode.com/GitHub_Trending/op/openpiJAX 训练出来的 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_SMUDGE1 uv sync # 拉全部依赖LeRobot 需跳过 LFS GIT_LFS_SKIP_SMUDGE1 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_dim、action_horizon、paligemma_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_dictq/k/v 全部 size mismatch模型一步都跑不起来。slice_gemma_state_dictpi0 与 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_typenp.ndarray, dtypefloat32)第 401 行附近走 JAX 模型的 restore 路径让 dtype 转换和 JAX 训练时一致。绕过它直接读原始 shard混合精度检查点转换后的推理数值会整体漂移。踩坑对照表⚠️ 只收录实际跑转换时的高频报错不凑数现象报错关键字根因一句话修复命令或代码≤ 2 行Invalid precision: float16脚本只落地支持 float32 / bfloat16float16 走 else 抛错改--precision bfloat16默认值可省略Error: --output_path is required没传--inspect_only却忘了输出路径补--output_path dirConfig xxx is not a Pi0Config--config_name指到了 pi0_fast 等非 flow 版本换成 pi0 / pi05 系列如pi0_droidMissing key(s)/size mismatchpi0 与 pi05 归一化层结构不同检查点和 config 版本没对上--config_name与版本成对pi0 对pi0_droidpi05 对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),仅供参考