3步把openpi的JAX检查点转成PyTorch:模型转换完整实战指南
【免费下载链接】openpi项目地址: https://gitcode.com/GitHub_Trending/op/openpi
RuntimeError: 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 转 PyTorch
openpi 是 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_SMUDGE=1 uv sync GIT_LFS_SKIP_SMUDGE=1 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.json(action_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: # AdaRMS:Dense_0/kernel 映射到 dense.weight,bias 映射到 dense.bias llm_input_layernorm_kernel = state_dict.pop( f"llm/layers/pre_attention_norm_{num_expert}/Dense_0/kernel{suffix}") else: # 普通 pi0:scale 向量直接映射到 layernorm.weight llm_input_layernorm = state_dict.pop( f"llm/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", framework="pt") 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.config=pi0_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),仅供参考