拓冰建站拓冰建站
首页 / 资讯中心 / 正文

3步把openpi的JAX检查点转成PyTorch:模型转换完整实战指南

3步把openpi的JAX检查点转成PyTorch模型转换完整实战指南【免费下载链接】openpi项目地址: https://gitcode.com/GitHub_Trending/op/openpiRuntimeError: Error(s) in loading state_dict: size mismatch for self_attn.q_proj.weight: copying a param with shape torch.Size([2048, 2048]) from checkpoint where the shape is torch.Size([18, 2048, 2048])想把 openpi 仓库的 JAX 检查点直接喂给 PyTorch 加载大概率撞上这类报错JAX 用 einsum 把注意力权重存成了每层一张融合大矩阵而 PyTorch 的nn.Linear要求拆好的 q/k/v/o 四个独立投影。本项目的 convert_jax_model_to_pytorch.py 转换脚本专门解决这件事。读完你能掌握三个技能检查 JAX 检查点的参数结构、跑通 π₀/π₀.₅ 的 JAX 转 PyTorch 全流程、并用转换后的模型直接起推理服务。项目定位为什么需要 JAX 转 PyTorchopenpi 是 Physical Intelligence 团队开源的 VLA视觉-语言-动作模型体系主力实现是 JAX覆盖 π₀、π₀-FAST、π₀.₅ 三类模型2025 年 9 月起仓库同步提供 PyTorch 版模型实现已在 LIBERO 基准上验证过推理与微调。由于 JAX 检查点是 Orbax 分片目录格式、参数命名也是 JAX 风格PyTorch 生态无法直接消费所以中间需要一步显式转换。整条链路就三步恢复原始权重 → 按子网络做维度与键映射 → 载入 PyTorch 模型落盘。逻辑全部在 convert_jax_model_to_pytorch.py 里几百行值得通读一遍。快速上手转换流程跑通步骤1. 克隆仓库并安装依赖项目用 uv 管理依赖LeRobot 走 git 引入安装时要跳过 LFS 以免拉大文件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 .2. 打 transformers 补丁PyTorch 版依赖 transformers 4.53.2 的三处修正支持 AdaRMS、激活精度控制、KV cache 不更新可用仓库自带补丁文件uv pip show transformers # 确认版本是 4.53.2 cp -r ./src/openpi/models_pytorch/transformers_replace/* .venv/lib/python3.11/site-packages/transformers/⚠️ 注意uv 默认 hardlink 模式下这会影响缓存里的 transformers想还原需uv cache clean transformers。3. 下载 JAX 检查点检查点首次使用时会缓存到~/.cache/openpi用仓库自带工具拉取pi0_droid并打印本地路径uv run python -c from openpi.shared import download; print(download.maybe_download(gs://openpi-assets/checkpoints/pi0_droid))输出形如/home/user/.cache/openpi/openpi-assets/checkpoints/pi0_droid下一步的--checkpoint_dir就用它。4. 先检查再转换用--inspect_only打印检查点内全部参数键的层级结构确认无误后再动手uv run examples/convert_jax_model_to_pytorch.py \ --checkpoint_dir /home/$USER/.cache/openpi/openpi-assets/checkpoints/pi0_droid \ --inspect_only确认没问题后去掉该参数、补上输出路径执行转换默认精度 bfloat16与 JAX 推理一致uv run 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✅ 预期输出终端打印Model conversion completed successfully!./pi0_droid_pytorch下生成model.safetensors权重、config.jsonaction_dim、precision 等、assets/归一化统计等资源。机制拆解维度错位到底发生在哪einsum 注意力权重拆成 q/k/v/o 投影JAX 侧llm/layers/attn/q_einsum/w是[层数, hidden, heads*head_dim]的融合张量PyTorch 侧没有对应结构。slice_paligemma_state_dict 对每一层做 transpose reshape 展开q_proj_weight_reshaped ( llm_attention_q_einsum[i] .transpose(0, 2, 1) .reshape(config.text_config.num_attention_heads * config.text_config.head_dim, config.text_config.hidden_size) )本质是把融合矩阵还原成nn.Linear期望的[out_features, in_features]。K/V 则从kv_einsum按索引拆开[i, 0, 0]是 K[i, 1, 0]是 V。卷积同理patch embedding 的 kernel 是 JAX 顺序[H, W, C_in, C_out]一句transpose(3, 2, 0, 1)换成 PyTorch 的[C_out, C_in, H, W]。pi05 自适应归一化分支π₀.₅ 把普通 RMSNorm 换成了自适应归一化原本的单个 scale 向量变成了一个小 Dense 层的 kernel bias参数键完全不同。slice_gemma_state_dict 靠路径名区分if pi05 in checkpoint_dir: # AdaRMSDense_0/kernel 映射到 dense.weightbias 映射到 dense.bias llm_input_layernorm_kernel state_dict.pop( fllm/layers/pre_attention_norm_{num_expert}/Dense_0/kernel{suffix}) else: # 普通 pi0scale 向量直接映射到 layernorm.weight llm_input_layernorm state_dict.pop( fllm/layers/pre_attention_norm_{num_expert}/scale{suffix})这也是按目录名识别版本的写法意味着你传入的checkpoint_dir必须是真实下载路径不能把 pi05 目录随意改名否则键会缺。常见问题与应对ValueError: Config xxx is not a Pi0Config— 原因config_name传成了 π₀-FAST 的配置FAST 是自回归头暂不支持转换。修复改用流匹配配置如--config_name pi0_droid或--config_name pi05_droid。size mismatch for ...— 原因config 与检查点对不上hidden 维度、action_dim 不同。修复对照 README 里的检查点清单让 config_name 与目录一一对应例如pi05_droid目录配--config_name pi05_droid。推理时 AdaRMS 相关 KeyError 或行为异常— 原因transformers 补丁没打或版本不是 4.53.2。修复回到第 2 步补补丁若缓存污染导致不生效先uv cache clean transformers再重装重打。报文件找不到、检查点打不开— 原因--checkpoint_dir要指向包含params/的检查点目录而不是某个具体文件。修复直接使用maybe_download返回的路径别手工拼接。uv sync 依赖冲突— 原因旧 venv 缓存残留。修复删掉.venv重新uv sync仍不行先uv self update升级 uv 本身。验证确认 JAX 转 PyTorch 模型可用先确认产物齐全ls ./pi0_droid_pytorch # 预期: model.safetensors config.json assets再用代码回读权重确认 safetensors 可加载、精度字段正确import json import safetensors with safetensors.safe_open(pi0_droid_pytorch/model.safetensors, frameworkpt) as f: keys list(f.keys()) print(len(keys)) # 数千个参数加载无异常 print(json.load(open(pi0_droid_pytorch/config.json))) # precision 应为 bfloat16最后起推理服务——PyTorch 版 API 与 JAX 完全一致create_trained_policy会自动识别检查点格式uv run scripts/serve_policy.py policy:checkpoint \ --policy.configpi0_droid --policy.dir./pi0_droid_pytorch服务能正常启动并用随机观测应答无机器人场景可参考examples/simple_client/的客户端示例说明整条链路已跑通。收尾openpi 转换工具的价值在于把einsum 权重拆不开、归一化键随版本变形、Orbax 无法直读这三个硬骨头压缩成一条命令。下一步可以尝试用scripts/train_pytorch.py直接微调转换后的模型在配置里把pytorch_weight_path指向转换输出或者阅读docs/remote_inference.md了解把策略服务部署到独立推理机的远程推理模式想参与改进可先看CONTRIBUTING.md。【免费下载链接】openpi项目地址: https://gitcode.com/GitHub_Trending/op/openpi创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门