十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

openpi:一条命令完成 JAX 转 PyTorch,pi0 checkpoint 导出 safetensors

openpi:一条命令完成 JAX 转 PyTorch,pi0 checkpoint 导出 safetensors openpi一条命令完成 JAX 转 PyTorchpi0 checkpoint 导出 safetensors【免费下载链接】openpi项目地址: https://gitcode.com/GitHub_Trending/op/openpi场景切入openpi 的 JAX 转 PyTorch 模型转换脚本就是为这类现场准备的仿真环境里 JAX 推理一切正常切到产线的 PyTorch 推理服务后load_state_dict直接抛 KeyError——orbax checkpoint 的参数键名和PI0Pytorch的 state_dict 完全对不上。examples/convert_jax_model_to_pytorch.py 一次导出即可不用逐 key 手动对维度。一次跑通仓库根目录跑uv sync装齐依赖JAX、orbax、torch 2.7.1、transformers 4.53.2 都在顶层 pyproject 里锁定PyTorch 侧无需另装。以 pi0_droid 为例先看参数结构再执行转换uv sync python examples/convert_jax_model_to_pytorch.py --checkpoint_dir /home/$USER/.cache/openpi/openpi-assets/checkpoints/pi0_droid --config_name pi0_droid --inspect_only python 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第一条只打印层级参数键树如llm/layers/attn/q_einsum/w先确认 checkpoint 完整--config_name按 checkpoint 目录名选pi0_droid、pi05_droid、pi0_aloha_sim均可第二条默认以 bfloat16 写出转换成功后输出目录有三样东西model.safetensors权重、config.json配置、assets/从 checkpoint 同级目录拷入的资源脚本内部在做什么错位的根源是两个框架存参数的方式不同JAXflax/nnx参数是嵌套 PyTree键为斜杠路径Linear 的 kernel 存[in_features, out_features]PyTorch 参数是module.attr平铺 dictnn.Linear的 weight 是[out_features, in_features]。脚本在三处做了处理卷积核维度转置PyTorch 卷积类权重要求[Cout, Cin, H, W]视觉塔的 patch embedding 要做四维重排直接照搬会 size mismatch# JAX 为 [H, W, Cin, Cout]PyTorch 需 [Cout, Cin, H, W] state_dict[pytorch_key] state_dict.pop(jax_key).transpose(3, 2, 0, 1)pi05 自适应归一化层 Dense 分支pi0 的归一化是普通 RMSNorm一维scale参数pi05 的动作专家改用 adaRMSNorm归一化层带Dense线性层kernel/bias两代版本键名不同需要按 checkpoint 目录名分支取参if pi05 in checkpoint_dir: llm_input_layernorm_kernel state_dict.pop(fllm/layers/pre_attention_norm_{num_expert}/Dense_0/kernel{suffix}) else: llm_input_layernorm state_dict.pop(fllm/layers/pre_attention_norm_{num_expert}/scale{suffix})MoE 多专家 state_dict 拆分与映射PaliGemma 基座与动作专家两套权重混在同一个 dict 里只靠键名上的_1后缀区分。脚本按expert_keys清单拆成两份映射再分别灌入paligemma和gemma_expert两个子模块for key, value in state_dict.items(): if key not in expert_keys: final_state_dict[key] torch.from_numpy(value) else: expert_dict[key] value出错速查 三个高频报错对照报错特征原因一句话修复命令或参数size mismatch for ...维度顺序错位或 config 与 checkpoint 不对应先--inspect_only核对维度再核对--config_name推理输出明显偏离 JAX 侧精度漂移--precision bfloat16Missing key(s) in state_dictconfig_name 与 checkpoint 模型版本不对应--config_name与 checkpoint 目录名核对如 pi05_droid验证与下一步回载校验加载后打印动作投影层形状与 config 的action_dim一致即权重完整——import openpi.training.config as _config, safetensors.torch from openpi.models_pytorch.pi0_pytorch import PI0Pytorch m PI0Pytorch(_config.get_config(pi0_droid).model) safetensors.torch.load_model(m, pi0_droid_pytorch/model.safetensors) print(m.action_out_proj.weight.shape) # 应与 config 的 action_dim 吻合确认形状与精度后即可接入 PyTorch 推理服务远程推理的完整部署流程见 docs/remote_inference.md转换中遇到的其他报错可直接提交 GitHub Issue。【免费下载链接】openpi项目地址: https://gitcode.com/GitHub_Trending/op/openpi创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表