一条指令跑通 JAX 转 PyTorch:openpi 权重迁移实战
【免费下载链接】openpi项目地址: https://gitcode.com/GitHub_Trending/op/openpi
你做 JAX 转 PyTorch 迁移,第一步load_state_dict直接抛size mismatch,注意力权重是三维的 einsum 形状,怎么对都对不齐。下面基于 openpi 的examples/convert_jax_model_to_pytorch.py脚本,把 pi0/pi05 的权重迁移完整跑一遍,最终产出一个可部署的model.safetensors,遇到不匹配也能直接定位到具体参数。
JAX 转 PyTorch 迁移时,权重经历了哪些变换
完整转换链如下:脚本读取 Orbax 检查点,按模块重排权重,最后写出 safetensors 转换产物。
JAX 与 PyTorch 的权重对不上,根因是三类结构差异:JAX 卷积权重是[H, W, C_in, C_out]顺序,PyTorch 是[C_out, C_in, H, W];JAX 版 Gemma 的 Q/O 投影存成四维 einsum 权重,PyTorch 的 Linear 只收二维矩阵;JAX 参数键是扁平字符串树,PyTorch 是点号模块路径。convert_pi0_checkpoint先经restore_params把 Orbax 检查点还原成 float32 的 numpy 字典,还原过程会尊重 JAX 模型加载时的 dtype 转换,数值不漂移,再交给两个 slice 函数逐模块重排。
# slice_paligemma_state_dict:重排 ViT 卷积核 state_dict[pytorch_key] = state_dict.pop(jax_key).transpose(3, 2, 0, 1)这一行把 JAX 的[H, W, C_in, C_out]转成 PyTorch 的[C_out, C_in, H, W]。
# 每层 q 投影:einsum 四维权重 -> Linear 二维矩阵 q_proj_weight_reshaped = ( llm_attention_q_einsum[i].transpose(0, 2, 1) .reshape(num_attention_heads * head_dim, hidden_size) )这一步把 einsum 的四维注意力权重压成nn.Linear可用的二维(out_features, in_features)矩阵。
# slice_gemma_state_dict:pi05 是自适应归一化,pi0 是 RMSNorm if "pi05" in checkpoint_dir: state_dict[f"{layer}.dense.weight"] = kernel.transpose() else: state_dict[f"{layer}.weight"] = scalepi05 把每层的 RMSNorm scale 参数换成了自适应归一化的 Dense 层,脚本靠检查点路径里是否带 pi05 区分这两种结构,写入不同键名。
ViT/LLM、动作专家、投影层三组权重最终合并成一个字典,装入新实例化的PI0Pytorch,按目标精度转换后经safetensors.torch.save_model写出model.safetensors,同时落一份config.json记录 action_dim、精度等字段,assets/目录也会一并拷贝过去。
⚡ 转换脚本实操路径:从检查点到 safetensors
开始前确认你在仓库根目录,后面所有命令都依赖uv run定位虚拟环境。
1. 同步依赖并给 transformers 打补丁
PyTorch 实现依赖 transformers 的三个补丁(AdaRMS、激活精度控制、KV 缓存复用),克隆后uv sync,再把补丁文件拷进 venv:
git clone https://gitcode.com/GitHub_Trending/op/openpi && cd openpi && uv sync && cp -r ./src/openpi/models_pytorch/transformers_replace/* .venv/lib/python3.11/site-packages/transformers/运行后你会看到:
Resolved 128 packages in 3s Installed 128 packages in 12s如果这里卡住:cp -r报 "No such file" 时,先跑uv pip show transformers确认 4.53.2 已装好再重试。 产物:带 .venv、且 transformers 已打补丁的仓库目录,后面所有命令都在此目录下执行。
2. 拉取 JAX 检查点并检查参数键结构
先用--inspect_only确认检查点键结构与脚本预期一致,避免转换白跑:
uv run examples/convert_jax_model_to_pytorch.py --checkpoint_dir ~/.cache/openpi/openpi-assets/checkpoints/pi05_droid --config_name pi05_droid --inspect_only输出是层级化的参数键列表:
img/embedding/kernel img/pos_embedding llm/embedder/input_embedding llm/layers/attn/q_einsum/w如果这里卡住:检查点目录不存在时,先运行uv run scripts/serve_policy.py droid让它从 gs://openpi-assets 下载完,Ctrl+C 后回来执行本命令。 产物:确认参数树包含 img/(视觉塔)、llm/(语言模型与动作专家)、投影层三组键,下一步复用同一路径。
3. 执行转换并输出 safetensors 权重
这条命令执行 JAX 转 PyTorch 的完整转换:以 float32 还原保证数值,输出精度默认 bfloat16,与 JAX 推理精度对齐:
uv run examples/convert_jax_model_to_pytorch.py --checkpoint_dir ~/.cache/openpi/openpi-assets/checkpoints/pi05_droid --config_name pi05_droid --output_path ./pi05_droid_pytorch运行后你会看到:
Converting PI0 checkpoint from ... to ./pi05_droid_pytorch Model conversion completed successfully! Model saved to ./pi05_droid_pytorch如果这里卡住:出现size mismatch时核对--config_name与检查点是否匹配,pi0 和 pi05 的投影层结构不同,交叉使用必报此错。 产物:./pi05_droid_pytorch/ 下三个文件——model.safetensors(权重)、config.json(记录 action_dim、精度等)、assets/(归一化统计),下一步把策略服务指向该目录。
4. 让策略服务指向转换后的模型
create_trained_policy会自动识别 PyTorch 格式,API 与 JAX 版一致,只换检查点路径:
uv run scripts/serve_policy.py policy:checkpoint --policy.config=pi05_droid --policy.dir=./pi05_droid_pytorch运行后你会看到:
INFO:root:Creating server (host: my-host, ip: 127.0.0.1)该行打印后,服务在 8000 端口监听,等待观测请求。 如果这里卡住:服务启动报 transformers 相关属性缺失,回到第 1 节重跑cp -r补丁命令。 产物:可调用的策略服务,客户端怎么接入看远程推理文档。
🔍 报错关键字排障速查与延伸入口
| 报错关键字 | 根因 | 修复动作 |
|---|---|---|
size mismatch for ... | --config_name与检查点不匹配 | 用匹配的 config_name 重跑第 3 节命令 |
Invalid precision | --precision不支持 float16 | 改用 bfloat16 或 float32 |
ModuleNotFoundError: openpi | 未在仓库内以 uv 方式运行 | 在仓库根目录用uv run ...执行命令 |
| transformers 缺 AdaRMS 相关属性 | transformers_replace 补丁未生效 | 重跑第 1 节的cp -r命令 |
gs://openpi-assets下载失败 | GCS 网络不可达 | 手动下载检查点到本地目录,--checkpoint_dir传本地路径 |
转换报错不在上表、或要继续 openpi 模型迁移链路时,去这三个地方:
- CONTRIBUTING 贡献指南:按模板提交转换报错 issue
- 规范统计文档:迁移后 state/action 归一化对不上时,检查 norm_stats 的处理
- 用下面 3 行代码自测转换后的检查点能否被策略正确加载:
from openpi.training import config as _config from openpi.policies import policy_config policy = policy_config.create_trained_policy(_config.get_config("pi05_droid"), "./pi05_droid_pytorch")【免费下载链接】openpi项目地址: https://gitcode.com/GitHub_Trending/op/openpi
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考