news 2026/9/12 5:57:25

一条指令跑通 JAX 转 PyTorch:openpi 权重迁移实战

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
一条指令跑通 JAX 转 PyTorch:openpi 权重迁移实战

一条指令跑通 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"] = scale

pi05 把每层的 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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/12 5:57:08

Windows部署vLLM指南:WSL2下运行Qwen3-8B-FP8推理服务

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/12 5:57:05

爬虫部署六大方案:从单机到分布式实战解析

1. 爬虫部署的核心挑战与解决思路在数据驱动的互联网时代,爬虫技术已经成为企业获取竞争情报、市场分析和业务决策的重要工具。但许多开发者都会遇到这样的困境:本地调试成功的爬虫脚本,一旦部署到生产环境就频繁崩溃。我曾为某电商平台部署价…

作者头像 李华
网站建设 2026/9/12 5:56:32

工业级提示词工程:Prompt as Code实践指南

1. 这不是又一个“AI画图工具”,而是一套工业级提示词交付系统你有没有遇到过这样的场景:团队里美术同学反复找你改提示词——“再加点赛博朋克感”“光影太硬,要柔焦”“主角衣服颜色偏暖一点”;开发同学在调试图像生成接口时&am…

作者头像 李华
网站建设 2026/9/12 5:55:22

Simulink建模助手:自动创建Bus Creator提升效率

1. 项目概述在Simulink建模过程中,信号线的管理往往成为影响工作效率的关键因素。当我们需要将多个From模块的输出信号汇总到一个Bus Creator时,传统的手动连线方式不仅耗时耗力,还容易出错。这个"根据From创建Bus Creator"的建模助…

作者头像 李华
网站建设 2026/9/12 5:55:04

Camofox-Browser深度解析:从Gecko内核构建到指纹扰动隐私实践

Camofox-Browser 这个项目,我从它早期原型就开始关注,最近总算在主力环境里把完整编译流程和日常使用都跑通了一遍。简单说,它是一款以隐私保护为核心目标的浏览器,主打指纹扰动、跟踪拦截和本地化数据处理。跟市面上常见的“套壳…

作者头像 李华
网站建设 2026/9/12 5:53:51

刘海屏菜单栏太乱?用 Ice 快速清出空间

刘海屏菜单栏太乱?用 Ice 快速清出空间 【免费下载链接】Ice Powerful menu bar manager for macOS 项目地址: https://gitcode.com/GitHub_Trending/ice/Ice Ice 管的是 macOS 菜单栏:图标能一键隐藏、拖拽排序、按需展开,还能换色加…

作者头像 李华