恒美微站 Logo 恒美微站
  • 首页
  • 关于我们
  • 建站服务
  • 主题模板
  • 案例展示
  • 资讯中心
  • 联系我们

两条命令把 JAX 权重搬进 PyTorch:openpi 模型转换实操

  • 首页
  • 资讯中心
  • /
  • 两条命令把 JAX 权重搬进 PyTorch:openpi 模型转换实操

相关资讯

Linux内核补丁机制与静态调用技术解析 2026/9/12 4:14:09
Haystack 的 PyversityRanker:用 pyversity 多样化算法在检索结果中平衡相关性与多样性 2026/9/12 4:14:09
AI应用开发实战:从可部署API到生产级AI工具的工程化路径 2026/9/12 4:14:09

最新资讯

Spring Boot集成OpenTelemetry实现分布式链路追踪实战
uutils coreutils 中 shuf 的基准测试指南:方法、命令与底层实现剖析
无人机三维路径规划:多目标遗传算法MATLAB实现
深入解析JavaScript闭包:原理与应用
SpringBoot美食分享系统开发实战
ASP.NET Core Razor多页签组件设计与业务整顿实践

今日推荐

MATLAB仿生优化框架:长鼻浣熊算法多策略融合实现
【JAVA毕设源码分享】基于 JavaWeb 的校园一卡通管理系统的设计与实现 基于 JavaWeb 的校园卡业务管理系统(程序+文档+代码讲解+一条龙定制)
【JAVA毕设源码分享】基于 Java 的图书馆借阅管理平台的搭建与实现 基于 Java 的图书馆综合管理系统(程序+文档+代码讲解+一条龙定制)

本周热门

超人会飞不算本事:系统稳定依赖清晰规则与边界设计
超人VS蜘蛛侠:拆解超级IP的影响力与传播方法论
基于CNN的调制信号识别:MATLAB实现时频图分类实战

本月精选

自研推理加速器Redwood:两周内实现PyTorch模型高效部署的实战教程
V4L2摄像头采集实战:从camera_client.rar到出图全流程解析
从“谁发明了钢琴键”到知识问答智能体:RAG与记忆工程实践

两条命令把 JAX 权重搬进 PyTorch:openpi 模型转换实操

发布时间:2026/9/12 4:14:09
两条命令把 JAX 权重搬进 PyTorch:openpi 模型转换实操 两条命令把 JAX 权重搬进 PyTorchopenpi 模型转换实操【免费下载链接】openpi项目地址: https://gitcode.com/GitHub_Trending/op/openpi某次JAX 侧训好的 π₀ 推理服务要并进一条纯 PyTorch 的部署流水线卡点只有一个权重得先转成 PyTorch 能加载的 safetensors。openpi 仓库自带的转换脚本把这件事压成先 inspect 参数树、再执行转换两条命令。30 秒看懂转换流程最反直觉的一环JAX 侧专家注意力是一整块 einsum 权重矩阵PyTorch 侧要拆成独立的 Q/K/V/O 投影矩阵脚本在中间替你做完 transpose 和 reshape 对齐。环境一行到位补丁别忘环境准备就一条git clone --recurse-submodules https://gitcode.com/GitHub_Trending/op/openpi进目录uv sync依赖用 uv 管理。脚本内部干的事读 OrbaxJAX 侧的检查点序列化格式里的参数字典按 PyTorch 结构重排键名、把权重转置成目标布局实例化PI0Pytorch后存成 safetensors。PyTorch 侧专家层用自适应 LayerNorm必须 patch 一下 transformers。uv sync cp -r ./src/openpi/models_pytorch/transformers_replace/* .venv/lib/python3.11/site-packages/transformers/看到cp无报错即可。这个 patch 的副作用见后文坑二提前知道少排查半天。先跑 inspect 确认参数树转换前先用--inspect_only把参数键连同 shape 和 dtype 打印出来uv run examples/convert_jax_model_to_pytorch.py \ --checkpoint_dir ~/.cache/openpi/openpi-assets/checkpoints/pi0_droid \ --config_name pi0_droid \ --inspect_only检查点默认缓存在~/.cache/openpi/openpi-assets/checkpoints/名字没有就先随便跑一次推理触发下载。看到层级键树PaliGemma/img/embedding/kernel、llm/layers/attn/q_einsum/w每行带 shapedtype即成功键名全带/value后缀说明是训练中间态检查点也属正常脚本能处理。--inspect_only要求目录下存在params/子目录报缺目录时先检查路径指到了检查点根目录而不是上层。执行转换并确认输出完整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 \ --precision bfloat16看到Model conversion completed successfully!即成功。--precision缺省就是 bfloat16与 JAX 侧推理精度一致可省略--config_name是必填项决定动作维度和模型结构。输出目录应长这样ls -lh ~/.cache/openpi/openpi-assets/checkpoints/pi0_droid_pytorch预期model.safetensorsbfloat16 下约 3GBconfig.json记录 precisionassets/归一化统计。assets由脚本从检查点同级目录复制缺了它推理会直接加载失败看到没有 assets 就先补。用 PyTorch 侧推理验一遍from openpi.training import config as _config from openpi.policies import policy_config config _config.get_config(pi0_droid) checkpoint_dir ~/.cache/openpi/openpi-assets/checkpoints/pi0_droid_pytorch # 换成实际路径 policy policy_config.create_trained_policy(config, checkpoint_dir) # 按 model.safetensors 自动识别 actions policy.infer(example)[actions] print(actions.shape)example是包含相机图像和prompt字段的观测字典droid 观测的完整键见 policy_config.py。输出 shape 为(8, 16)droid8 维动作 × 16 步预测即验证通过。路径不含 pi05 子串模型会静默降配当转换的模型是 pi05 系列、但检查点或输出目录名里没有pi05这个子串时通常是脚本走错了分支它用路径判断是否为 pi05见 convert_jax_model_to_pytorch.py 的pi05 in checkpoint_dir判错后专家层的自适应 LayerNorm 权重被当普通 RMSNorm 写入load_state_dict又是strictFalse缺键被静默跳过推理输出直接 NaN。if pi05 in checkpoint_dir: # slice_gemma_state_dict 内决定 dense 层还是 scale 层 state_dict[f{layer}.input_layernorm.dense.weight] kernel.transpose() else: state_dict[f{layer}.input_layernorm.weight] scale修复就一行约定pi05 检查点的路径里保留pi05子串跑完再抽查 state_dict 键名是否带.dense即可确认没走偏。uv 缓存里的 transformers 补丁会串门当你在多个项目共用 uv 时上面那条cp补丁会写进 uv 缓存——默认硬链接模式下改动永久生效重装 transformers 都带不回来别的项目里 transformers 行为悄悄变了。uv cache clean transformers # 彻底还原精度参数只认两种值当--precision传了float16时通常是踩了签名与实现不一致参数类型标注是float32 | bfloat16 | float16函数体只处理前两种第三个值直接ValueError。if precision float32: pi0_model pi0_model.to(torch.float32) elif precision bfloat16: pi0_model pi0_model.to(torch.bfloat16) else: raise ValueError(fInvalid precision: {precision})推理场景保持默认 bfloat16 即可与 JAX 侧推理精度一致文件体积比 float32 省一半。上线前对照一遍输出目录三件套齐全model.safetensors、config.jsonprecision 字段与预期一致、assets/归一化统计在--config_name与检查点家族匹配pi0_droid对 pi0pi05_droid对 pi05写错会先撞动作维度 shape 错误pi05 检查点路径含pi05子串且 state_dict 键名带.denseCPU 上验证时留意bfloat16 权重在 CPU 会被自动降到 float32输出数值与 GPU 可能有微小偏差迁移线上时把assets/随检查点目录一起搬不要只拷model.safetensors转换细节看 examples/convert_jax_model_to_pytorch.pyconfig 定义在 src/openpi/training/config.py转换完直接接 scripts/train_pytorch.py 就能在 PyTorch 侧继续微调。【免费下载链接】openpi项目地址: https://gitcode.com/GitHub_Trending/op/openpi创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

关于恒美微站

恒美微站专注于为个体商户、工作室提供极简自助建站服务,让每个人都能轻松拥有专业网站。

快速链接

  • 关于我们
  • 建站服务
  • 主题模板
  • 案例展示
  • 资讯中心

服务项目

  • 可视化建站
  • 拖拽编辑
  • 主题定制
  • SEO 优化
  • 网站托管

联系方式

  • 📍 地址:北京市朝阳区建国路 88 号
  • 📞 电话:400-888-8888
  • ✉️ 邮箱:info@hmyw.cn
  • 🕐 时间:周一至周日 9:00-18:00

© 2024 恒美微站 hmyw.cn 版权所有 | 京 ICP 备 12345678 号